Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 35 additions & 3 deletions powersig/jax/jax_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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'):
Expand Down
3 changes: 2 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -43,14 +43,15 @@ 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",
]
all = [
"jax[cuda13]>=0.10.0",
"torch>=2.5.0",
"cupy-cuda12x>=13.4.1",
"cupy-cuda13x>=13.6.0",
]

[project.urls]
Expand Down
3 changes: 3 additions & 0 deletions requirements-dev.txt
Original file line number Diff line number Diff line change
@@ -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
Expand Down
15 changes: 11 additions & 4 deletions requirements.txt
Original file line number Diff line number Diff line change
@@ -1,17 +1,24 @@
# 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
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
Loading