From c511158ee2822ea64e8169e981e664dc6a9e65ea Mon Sep 17 00:00:00 2001 From: Matthew Tamayo-Rios Date: Thu, 20 Aug 2026 19:44:43 +0000 Subject: [PATCH 1/2] Move the GPU stack to CUDA 13 and numpy 2 MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit requirements.txt asked for cupy-cuda12x and jax[cuda12] while pyproject's jax-gpu and all extras already asked for jax[cuda13], so the two disagreed about which CUDA major this project targets. Both cupy wheels provide the same `cupy` module, so installing across the two leaves one of them shadowed rather than raising a conflict. The numpy pin is the part that quietly costs a GPU. jax only gained CUDA 13 plugins in 0.10, and jax>=0.10 requires numpy>=2.0, so numpy==1.26.4 forces the resolver down to a jax with no CUDA 13 plugin at all. It installs, imports, and runs on CPU next to a GPU-enabled torch and cupy, warning only in a log line. Relaxing to numpy>=2.1 pulls jax-cuda13-plugin and jax-cuda13-pjrt in. The <2.3 ceiling is numba 0.61.2's, so that pin stays as is. torch's floor moves from 2.2.2 to 2.5.0 to match the one pyproject already declares; CUDA 13 additionally needs a cu130 build, which current PyPI wheels provide. pyproject's `all` extra combined jax[cuda13] with cupy-cuda12x, which cannot both be right. It now uses cupy-cuda13x. The existing `cupy` extra is left on cuda12x for anyone still on that toolkit, with `cupy-cuda13` added alongside. Verified on CUDA 13.0 (driver 580.178.04, 4x GPU): torch 2.13.0+cu130, cupy-cuda13x, and jax 0.11.1 all report GPU; numba 0.61.2 compiles against numpy 2.2.6. Dependency metadata only — no source changes. --- pyproject.toml | 3 ++- requirements-dev.txt | 3 +++ requirements.txt | 15 +++++++++++---- 3 files changed, 16 insertions(+), 5 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 378e213..34fc020 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -43,6 +43,7 @@ jax-cpu = ["jax[cpu]>=0.4.34"] jax-gpu = ["jax[cuda13]>=0.10.0"] torch = ["torch>=2.5.0"] cupy = ["cupy-cuda12x>=13.4.1"] +cupy-cuda13 = ["cupy-cuda13x>=13.6.0"] dev = [ "pytest>=7.0", "fbm>=0.3.0", @@ -50,7 +51,7 @@ dev = [ all = [ "jax[cuda13]>=0.10.0", "torch>=2.5.0", - "cupy-cuda12x>=13.4.1", + "cupy-cuda13x>=13.6.0", ] [project.urls] diff --git a/requirements-dev.txt b/requirements-dev.txt index 8446304..029be2d 100644 --- a/requirements-dev.txt +++ b/requirements-dev.txt @@ -1,4 +1,7 @@ matplotlib +# sigkernel's setup.py imports Cython without declaring it in +# [build-system].requires, so a PEP 517 isolated build fails. Install it with +# `pip install --no-build-isolation` after the rest of this file. git+https://github.com/crispitagorico/sigkernel.git git+https://github.com/geekbeast/KSig.git git+https://github.com/FrancescoPiatti/polysigkernel.git diff --git a/requirements.txt b/requirements.txt index 9e8b7c1..4e73976 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,7 +1,14 @@ +# CUDA 13 stack. cupy-cuda12x and cupy-cuda13x both provide the `cupy` module, +# so exactly one of them belongs here. torch needs a cu130 build (2.9+) to match; +# the current default PyPI wheel provides one. +# +# numpy is >=2.1 because jax only gained CUDA 13 plugins in 0.10, and jax>=0.10 +# requires numpy>=2.0 — pinning numpy to 1.26 silently leaves jax on CPU. The +# <2.3 ceiling is numba 0.61.2's. Cython==3.0.11 -numpy==1.26.4 +numpy>=2.1,<2.3 scipy -torch>=2.2.2 +torch>=2.5.0 torchvision torchaudio scikit-learn @@ -9,9 +16,9 @@ wheel setuptools aeon numba==0.61.2 -cupy-cuda12x==13.5.1 +cupy-cuda13x>=13.6.0 psutil tqdm==4.67.1 -jax[cuda12] +jax[cuda13]>=0.10.0 yfinance matplotlib From 01f9deb10b8f91be243e09b4c785fedeb1ce47dd Mon Sep 17 00:00:00 2001 From: Matthew Tamayo-Rios Date: Thu, 20 Aug 2026 19:55:10 +0000 Subject: [PATCH 2/2] Keep configure_jax working across the JAX 0.10/0.11 config rename MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit JAX 0.11 removed jax_exec_time_optimization_effort and jax_memory_fitting_effort, replacing the continuous 0.0-1.0 floats with the jax_optimization_level and jax_memory_fitting_level enums (UNKNOWN/O0-O3). jax.config.update raises AttributeError on an unrecognised key, so on 0.11 configure_jax died partway through — after enabling x64 and setting the XLA flags, before the matmul precision and the compilation cache were ever configured. This project declares jax>=0.10.0 and 0.10 only accepts the old spelling, so pinning to <0.11 would trade one broken half of the declared range for the other. _update_first_supported instead applies the first key the installed JAX recognises, leaving both ends of the range working. Values carry the old intent over: 1.0 was maximum execution-time effort, so O3. Memory fitting was deliberately set to 0.3 for high-memory machines and the new default is O2, so O1 keeps it below default. Verified against jax 0.11.1 on 4x H100 — configure_jax reports O3/O1 with x64 live and GPU detected, and tests/test_core_jax.py passes 21/21. --- powersig/jax/jax_config.py | 38 +++++++++++++++++++++++++++++++++++--- 1 file changed, 35 insertions(+), 3 deletions(-) diff --git a/powersig/jax/jax_config.py b/powersig/jax/jax_config.py index d592b28..451e4a0 100644 --- a/powersig/jax/jax_config.py +++ b/powersig/jax/jax_config.py @@ -10,6 +10,24 @@ CPU_COUNT = 32 # High core count TOTAL_MEMORY_GB = 64 # High memory (64GB) +def _update_first_supported(candidates, what): + """Apply the first config key the installed JAX actually recognises. + + jax.config keys are not stable across minor releases and jax.config.update + raises AttributeError on an unknown one, which aborts the rest of this + function. JAX 0.11 removed the continuous `*_effort` floats in favour of + O0-O3 enums; PowerSig declares jax>=0.10.0, so both spellings have to work. + """ + for key, value in candidates: + try: + jax.config.update(key, value) + return key + except (AttributeError, ValueError): + continue + print(f"JAX {jax.__version__}: no supported config key for {what}, using the default") + return None + + def configure_jax(): # Enable 64-bit precision jax.config.update('jax_enable_x64', True) @@ -26,15 +44,29 @@ def configure_jax(): # Enable optimizations for speed jax.config.update('jax_disable_most_optimizations', False) - jax.config.update('jax_exec_time_optimization_effort', 1.0) + # Maximum execution-time optimization. JAX 0.11 replaced the float + # jax_exec_time_optimization_effort (0.0-1.0) with the jax_optimization_level + # enum (O0-O3); 1.0 was the maximum, so O3. + _update_first_supported( + [('jax_optimization_level', 'O3'), + ('jax_exec_time_optimization_effort', 1.0)], + 'execution-time optimization', + ) jax.config.update('jax_default_matmul_precision', 'highest') # Enable and configure compilation cache jax.config.update('jax_enable_compilation_cache', True) jax.config.update('jax_compilation_cache_max_size', 2048 * 1024 * 1024) # 2GB cache - # Set memory fitting effort for high-memory systems - jax.config.update('jax_memory_fitting_effort', 0.3) + # Set memory fitting effort for high-memory systems. Same rename: + # jax_memory_fitting_effort -> jax_memory_fitting_level. 0.3 was deliberately + # low ("plenty of RAM, don't burn compile time squeezing"), and the new + # default is O2, so O1 keeps it below default. + _update_first_supported( + [('jax_memory_fitting_level', 'O1'), + ('jax_memory_fitting_effort', 0.3)], + 'memory fitting', + ) # Set persistent cache directory if not os.path.exists('/tmp/jax_cache'):