From d92a9143fe97c4e05b2d479af25a553636c1da85 Mon Sep 17 00:00:00 2001 From: Matthew Tamayo Date: Thu, 2 Apr 2026 14:15:46 -0700 Subject: [PATCH 1/3] Switch to pyproject + enable CI. --- .github/workflows/ci.yml | 33 ++++++ pyproject.toml | 61 +++++++++++ setup.py | 38 ------- tests/test_core_jax.py | 226 +++++++++++++++++++++++++++++++++++++++ 4 files changed, 320 insertions(+), 38 deletions(-) create mode 100644 .github/workflows/ci.yml create mode 100644 pyproject.toml delete mode 100644 setup.py create mode 100644 tests/test_core_jax.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..77911fd --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,33 @@ +name: CI + +on: + push: + branches: [main] + pull_request: + branches: [main] + +jobs: + test: + runs-on: ubuntu-latest + strategy: + matrix: + python-version: ["3.12"] + + steps: + - uses: actions/checkout@v4 + + - name: Set up Python ${{ matrix.python-version }} + uses: actions/setup-python@v5 + with: + python-version: ${{ matrix.python-version }} + + - name: Install dependencies + run: | + python -m pip install --upgrade pip + pip install ".[jax-cpu,dev]" + + - name: Run tests + env: + JAX_PLATFORMS: cpu + run: | + pytest tests/test_core_jax.py -v diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..d2e6752 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,61 @@ +[build-system] +requires = ["setuptools>=68.0"] +build-backend = "setuptools.build_meta" + +[project] +name = "powersig" +version = "1.0.0" +description = "Signature Kernel Power Series Library" +readme = "README.md" +license = "Apache-2.0" +requires-python = ">=3.12" +authors = [ + { name = "Matthew Tamayo-Rios", email = "matthew@geekbeast.com" }, +] +keywords = [ + "machine-learning", + "signature", + "sequence", + "time-series", + "deep-learning", +] +classifiers = [ + "Intended Audience :: Developers", + "Intended Audience :: Information Technology", + "Intended Audience :: Science/Research", + "Natural Language :: English", + "Operating System :: MacOS :: MacOS X", + "Operating System :: Microsoft :: Windows", + "Operating System :: Unix", + "Programming Language :: Python :: 3.12", + "Topic :: Scientific/Engineering :: Artificial Intelligence", + "Topic :: Scientific/Engineering :: Information Analysis", + "Topic :: Scientific/Engineering :: Mathematics", +] +dependencies = [ + "numpy>=1.26.4", + "scikit-learn>=1.3.2", + "tqdm>=4.67.1", +] + +[project.optional-dependencies] +jax-cpu = ["jax[cpu]>=0.4.0"] +jax-gpu = ["jax[cuda12]>=0.4.0"] +torch = ["torch>=2.5.0"] +cupy = ["cupy-cuda12x>=13.4.1"] +dev = [ + "pytest>=7.0", +] +all = [ + "jax[cuda12]>=0.4.0", + "torch>=2.5.0", + "cupy-cuda12x>=13.4.1", +] + +[project.urls] +Homepage = "https://github.com/geekbeast/powersig" +Repository = "https://github.com/geekbeast/powersig" + +[tool.setuptools.packages.find] +include = ["powersig*"] +exclude = ["tests*", "benchmarks*", "examples*"] diff --git a/setup.py b/setup.py deleted file mode 100644 index 342f9c1..0000000 --- a/setup.py +++ /dev/null @@ -1,38 +0,0 @@ -from setuptools import setup, find_packages - -def read_file(filename): - with open(filename, encoding='utf-8') as f: - return f.read().strip() - -version = read_file('VERSION') -readme = read_file('README.md') - -setup( - name='powersig', - version=version, - author='Matthew Tamayo-Rios', - author_email='matthew@geekbeast.com', - description='Signature Kernel Power Series Library', - long_description=readme, - long_description_content_type='text/markdown', - license='Apache License 2.0', - keywords='machine-learning signature sequence time-series pinn deep learning', - url='https://github.com/geekbeast/powersig', - packages=find_packages(), - install_requires=['torch>=2.5.0', 'numpy>=1.26.4', 'scikit-learn>=1.3.2', 'tqdm==4.67.1'], - python_requires='>=3.12', - classifiers=[ - 'Audience :: Developers', - 'Intended Audience :: Information Technology', - 'Intended Audience :: Science/Research', - 'License :: OSI Approved :: Apache Software License', - 'Natural Language :: English', - 'Operating System :: MacOS :: MacOS X', - 'Operating System :: Microsoft :: Windows', - 'Operating System :: Unix', - 'Programming Language :: Python :: 3.12', - 'Topic :: Scientific/Engineering :: Artificial Intelligence', - 'Topic :: Scientific/Engineering :: Information Analysis', - 'Topic :: Scientific/Engineering :: Mathematics', - ], -) diff --git a/tests/test_core_jax.py b/tests/test_core_jax.py new file mode 100644 index 0000000..4b5d7d7 --- /dev/null +++ b/tests/test_core_jax.py @@ -0,0 +1,226 @@ +"""CPU-compatible tests for the core JAX signature kernel implementation. + +These tests are designed to run without GPU, torch, cupy, ksig, or sigkernel +dependencies, making them suitable for CI on free-tier GitHub Actions runners. +""" + +import unittest +from math import ceil, sqrt + +import jax +import jax.numpy as jnp +import numpy as np + +from powersig.jax.algorithm import ( + PowerSigJax, + batch_ADM_for_diagonal, + build_stencil, + compute_block_size, + compute_vandermonde_vectors, + estimate_bytes_per_pair, + get_available_gpu_memory, + _round_to_power_of_2, +) +from powersig.util.grid import get_diagonal_range + + +# --------------------------------------------------------------------------- +# Stencil construction +# --------------------------------------------------------------------------- +class TestBuildStencil(unittest.TestCase): + def setUp(self): + self.order = 4 + self.dtype = jnp.float64 + + def test_shape(self): + stencil = build_stencil(self.order, self.dtype) + self.assertEqual(stencil.shape, (self.order, self.order)) + + def test_first_row_and_column_are_ones(self): + stencil = build_stencil(self.order, self.dtype) + np.testing.assert_allclose(np.array(stencil[0, :]), np.ones(self.order)) + np.testing.assert_allclose(np.array(stencil[:, 0]), np.ones(self.order)) + + def test_known_values(self): + stencil = build_stencil(self.order, self.dtype) + expected = np.array([ + [1.0, 1.0, 1.0, 1.0], + [1.0, 1.0, 0.5, 1 / 3], + [1.0, 0.5, 0.25, 1 / 12], + [1.0, 1 / 3, 1 / 12, 1 / 36], + ]) + np.testing.assert_allclose(np.array(stencil), expected, rtol=1e-10) + + def test_symmetry(self): + stencil = build_stencil(self.order, self.dtype) + np.testing.assert_allclose( + np.array(stencil), np.array(stencil.T), rtol=1e-10 + ) + + +# --------------------------------------------------------------------------- +# Vandermonde vectors +# --------------------------------------------------------------------------- +class TestVandermondeVectors(unittest.TestCase): + def test_unit_step(self): + v_s, v_t = compute_vandermonde_vectors(1.0, 1.0, 4, jnp.float64) + np.testing.assert_allclose(np.array(v_s), np.ones(4)) + np.testing.assert_allclose(np.array(v_t), np.ones(4)) + + def test_power_scaling(self): + v_s, v_t = compute_vandermonde_vectors(0.5, 0.25, 4, jnp.float64) + np.testing.assert_allclose( + np.array(v_s), [1.0, 0.5, 0.25, 0.125], rtol=1e-10 + ) + np.testing.assert_allclose( + np.array(v_t), [1.0, 0.25, 0.0625, 0.015625], rtol=1e-10 + ) + + +# --------------------------------------------------------------------------- +# Diagonal grid geometry +# --------------------------------------------------------------------------- +class TestDiagonalRange(unittest.TestCase): + def test_square_grid(self): + # 3x3 grid: diagonals 0,1,2 + s, t, dlen = get_diagonal_range(0, 3, 3) + self.assertEqual((s, t, dlen), (0, 0, 1)) + + s, t, dlen = get_diagonal_range(1, 3, 3) + self.assertEqual((s, t, dlen), (1, 0, 2)) + + s, t, dlen = get_diagonal_range(2, 3, 3) + self.assertEqual((s, t, dlen), (2, 0, 3)) + + def test_rectangular_grid(self): + # 2 rows, 4 cols + s, t, dlen = get_diagonal_range(0, 2, 4) + self.assertEqual((s, t, dlen), (0, 0, 1)) + + s, t, dlen = get_diagonal_range(3, 2, 4) + self.assertEqual((s, t, dlen), (3, 0, 2)) + + s, t, dlen = get_diagonal_range(4, 2, 4) + self.assertEqual((s, t, dlen), (3, 1, 1)) + + +# --------------------------------------------------------------------------- +# Block size auto-tuning utilities +# --------------------------------------------------------------------------- +class TestBlockSizeUtils(unittest.TestCase): + def test_round_to_power_of_2(self): + self.assertEqual(_round_to_power_of_2(1), 1) + self.assertEqual(_round_to_power_of_2(2), 2) + self.assertEqual(_round_to_power_of_2(3), 4) + self.assertEqual(_round_to_power_of_2(5), 8) + self.assertEqual(_round_to_power_of_2(16), 16) + self.assertEqual(_round_to_power_of_2(17), 32) + + def test_estimate_bytes_per_pair(self): + bpp = estimate_bytes_per_pair(100, 32, jnp.float64) + # 8 * 100 * (7*32 + 3*32^2) = 8 * 100 * 3296 = 2_636_800 + self.assertEqual(bpp, 2_636_800) + + def test_compute_block_size_bounded(self): + device = jax.devices("cpu")[0] + bs = compute_block_size(100, 32, jnp.float64, device, 1000) + self.assertGreaterEqual(bs, 1) + self.assertLessEqual(bs, 256) + # Must be a power of 2 + self.assertEqual(bs & (bs - 1), 0) + + def test_compute_block_size_clamped_to_total(self): + device = jax.devices("cpu")[0] + bs = compute_block_size(1, 4, jnp.float64, device, 3) + self.assertLessEqual(bs, 4) # rounded power of 2 of min(computed, 3) + + +# --------------------------------------------------------------------------- +# Gram matrix computation (end-to-end) +# --------------------------------------------------------------------------- +class TestGramMatrix(unittest.TestCase): + def setUp(self): + self.ps = PowerSigJax(order=8, device=jax.devices("cpu")[0]) + key = jax.random.PRNGKey(42) + self.X = jax.random.normal(key, (4, 10, 3)) + self.Y = jax.random.normal(jax.random.PRNGKey(99), (4, 10, 3)) + + def test_block_size_1_matches_auto(self): + gram_seq = self.ps.compute_gram_matrix(self.X, self.Y, block_size=1) + gram_auto = self.ps.compute_gram_matrix(self.X, self.Y) + np.testing.assert_allclose( + np.array(gram_seq), np.array(gram_auto), rtol=1e-10 + ) + + def test_explicit_block_sizes_match(self): + gram_1 = self.ps.compute_gram_matrix(self.X, self.Y, block_size=1) + gram_4 = self.ps.compute_gram_matrix(self.X, self.Y, block_size=4) + gram_16 = self.ps.compute_gram_matrix(self.X, self.Y, block_size=16) + np.testing.assert_allclose(np.array(gram_1), np.array(gram_4), rtol=1e-10) + np.testing.assert_allclose(np.array(gram_1), np.array(gram_16), rtol=1e-10) + + def test_symmetric(self): + gram = self.ps.compute_gram_matrix(self.X, self.X, symmetric=True) + np.testing.assert_allclose( + np.array(gram), np.array(gram.T), rtol=1e-10 + ) + + def test_symmetric_matches_full(self): + gram_full = self.ps.compute_gram_matrix(self.X, self.X, symmetric=False) + gram_sym = self.ps.compute_gram_matrix(self.X, self.X, symmetric=True) + np.testing.assert_allclose( + np.array(gram_full), np.array(gram_sym), rtol=1e-10 + ) + + def test_diagonal_positive(self): + """Signature kernel of a path with itself should be positive.""" + gram = self.ps.compute_gram_matrix(self.X, self.X, symmetric=True) + diag = np.diag(np.array(gram)) + self.assertTrue(np.all(diag > 0), f"Diagonal has non-positive entries: {diag}") + + def test_single_entry_matches_gram(self): + """compute_signature_kernel should match the corresponding Gram entry.""" + gram = self.ps.compute_gram_matrix(self.X, self.Y) + for i in range(min(2, self.X.shape[0])): + for j in range(min(2, self.Y.shape[0])): + single = self.ps.compute_signature_kernel(self.X[i], self.Y[j]) + np.testing.assert_allclose( + float(gram[i, j]), float(single), rtol=1e-6, + err_msg=f"Mismatch at ({i},{j})" + ) + + def test_call_interface(self): + """__call__ should produce the same result as compute_gram_matrix.""" + gram_method = self.ps.compute_gram_matrix(self.X, self.Y) + gram_call = self.ps(self.X, self.Y) + np.testing.assert_allclose( + np.array(gram_method), np.array(gram_call), rtol=1e-10 + ) + + def test_call_with_block_size(self): + gram = self.ps(self.X, self.Y, block_size=2) + gram_ref = self.ps(self.X, self.Y, block_size=1) + np.testing.assert_allclose( + np.array(gram), np.array(gram_ref), rtol=1e-10 + ) + + +# --------------------------------------------------------------------------- +# Batch ADM +# --------------------------------------------------------------------------- +class TestBatchADM(unittest.TestCase): + def test_2x2(self): + rho = jnp.array([0.5, 0.7], dtype=jnp.float64) + S = jnp.array([[10.0, 30.0], [100.0, 300.0]], dtype=jnp.float64) + T = jnp.array([[10.0, 20.0], [100.0, 200.0]], dtype=jnp.float64) + stencil = jnp.array([[1.0, 2.0], [3.0, 4.0]], dtype=jnp.float64) + U_buf = jnp.empty((2, 2, 2), dtype=jnp.float64) + + result = batch_ADM_for_diagonal(rho, U_buf, S, T, stencil) + self.assertEqual(result.shape, (2, 2, 2)) + # Verify not all zeros (computation happened) + self.assertFalse(jnp.allclose(result[:2], jnp.zeros_like(result[:2]))) + + +if __name__ == "__main__": + unittest.main() From f6cb748b617dc23df39b3ba437bd2e70aa493477 Mon Sep 17 00:00:00 2001 From: Matthew Tamayo Date: Thu, 2 Apr 2026 14:30:09 -0700 Subject: [PATCH 2/3] Add build status and versioning shields. --- README.md | 4 ++++ pyproject.toml | 5 ++++- 2 files changed, 8 insertions(+), 1 deletion(-) diff --git a/README.md b/README.md index 32db1d9..d449cfc 100644 --- a/README.md +++ b/README.md @@ -1,4 +1,8 @@ # Power-series based computation of Signature Kernels + +[![CI](https://github.com/geekbeast/powersig/actions/workflows/ci.yml/badge.svg)](https://github.com/geekbeast/powersig/actions/workflows/ci.yml) +[![Version](https://img.shields.io/badge/version-1.0.0-blue)](https://github.com/geekbeast/powersig) + Using ADM-derived Neumann series to compute signature kernels. ## Installation diff --git a/pyproject.toml b/pyproject.toml index d2e6752..11dc746 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "powersig" -version = "1.0.0" +dynamic = ["version"] description = "Signature Kernel Power Series Library" readme = "README.md" license = "Apache-2.0" @@ -56,6 +56,9 @@ all = [ Homepage = "https://github.com/geekbeast/powersig" Repository = "https://github.com/geekbeast/powersig" +[tool.setuptools.dynamic] +version = {file = "VERSION"} + [tool.setuptools.packages.find] include = ["powersig*"] exclude = ["tests*", "benchmarks*", "examples*"] From 8ac146d9cf13156cbc3811c1ed4823072ca2ec34 Mon Sep 17 00:00:00 2001 From: Matthew Tamayo Date: Thu, 2 Apr 2026 17:25:14 -0700 Subject: [PATCH 3/3] fix ci. --- powersig/__init__.py | 51 ++++++++++++++++++++++++++++---------------- 1 file changed, 33 insertions(+), 18 deletions(-) diff --git a/powersig/__init__.py b/powersig/__init__.py index 4d96b4c..d4957c4 100644 --- a/powersig/__init__.py +++ b/powersig/__init__.py @@ -2,28 +2,43 @@ PowerSig - Efficient Computation of Signature Kernels This package provides efficient implementations of signature kernels -using both JAX and CuPy for GPU acceleration. +using JAX, PyTorch, or CuPy backends. Install only the backend you need: + + pip install powersig[jax-cpu] # JAX on CPU + pip install powersig[jax-gpu] # JAX on GPU + pip install powersig[torch] # PyTorch + pip install powersig[cupy] # CuPy """ -# Import submodules first -from . import jax -from . import torch -from . import util -from . import cupy_backend +import importlib as _importlib + + +def __getattr__(name): + """Lazy-import backend submodules so missing optional deps don't break import.""" + _submodules = {"jax", "torch", "cupy_backend", "util"} + if name in _submodules: + return _importlib.import_module(f".{name}", __name__) + + # Convenience re-exports — only attempt if the backend is installed + _lazy_imports = { + "PowerSigJax": (".jax.algorithm", "PowerSigJax"), + "fractional_brownian_motion": (".jax.utils", "fractional_brownian_motion"), + "fbm": (".util.fbm_utils", "fractional_brownian_motion"), + } + if name in _lazy_imports: + module_path, attr = _lazy_imports[name] + mod = _importlib.import_module(module_path, __name__) + return getattr(mod, attr) -# Main implementations -from .jax.algorithm import PowerSigJax -from .jax.utils import fractional_brownian_motion + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") -# Utility functions -from .util.fbm_utils import fractional_brownian_motion as fbm __all__ = [ - 'PowerSigJax', - 'fractional_brownian_motion', - 'fbm', - 'jax', - 'torch', - 'util', - 'cupy_backend' + "PowerSigJax", + "fractional_brownian_motion", + "fbm", + "jax", + "torch", + "util", + "cupy_backend", ]