Safe persistent JAX compile cache, on by default; no AF3 recompile on the second prediction - #644
Merged
Merged
Conversation
With --jax_compilation_cache_dir, JAX also writes XLA's per-fusion autotune
cache into the directory. Inserting into it fails an atomic rename on shared
filesystems ("Failed to insert autotune cache: FAILED_PRECONDITION") and XLA
then segfaults while compiling: seen with the AF3 image's jaxlib 0.9.1 on
BeeGFS /scratch and on a job's TMPDIR. This is what makes AF3 batching crash
in AlphaPulldownSnakemake, which points the cache at the output directory.
Both backends now go through one helper that keeps only JAX's own entries
(jax_persistent_cache_enable_xla_caches=none, unless the environment sets
it), caches every executable, and skips the cache with a warning when the
directory cannot be created or written instead of failing the prediction.
AF2 (jax 0.5.3) carries the same per-fusion cache, so it gets the same fix.
Measured on 9 GPU types: an AF3 call served from the cache is 1.6-2.0x
faster on average (3-4.5x at 164 tokens), outputs bit-identical, and the
cache works across processes and across complexes in one token bucket.
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
An unset --jax_compilation_cache_dir now resolves to $JAX_COMPILATION_CACHE_DIR or, failing that, $XDG_CACHE_HOME/alphapulldown/jax_compilation_cache (~/.cache/...), so direct runs share compiled models across processes too. "none", "false" or "" turns the cache off; an explicit path is used as given. Programmatic setup(jax_compilation_cache_dir=None) keeps caching off. The per-user default is pruned oldest-first to 10 GB when a backend starts: AF3 adds ~3 MB per token bucket, but AF2 ~5 MB per new complex size without bound. Pruning is lock-free and best effort; a vanished entry only costs a recompile. A directory the user chose is never pruned. jax's own size limit needs filelock, which the images do not ship. README: the cache note now matches AlphaPulldownSnakemake's. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Merged
5 tasks done
|
You have reached your Codex usage limits for code reviews. You can see your limits in the Codex usage dashboard. |
…re, not per call Mirrors AlphaPulldownSnakemake's README. The module docstring no longer says AF3 recompiles on every prediction call: a process compiles each token bucket once, plus once more on its second prediction. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
tokamax creates a JAX user context (jax.make_user_context) the first time an op consults its autotuning cache, which happens while the model is first traced. A new user context adds an entry to JAX's trace context and so changes every jit cache key: the second predict call in a process missed the in-memory cache and re-traced and recompiled the whole model (~50 s), and only later calls hit. AlphaFold3Backend.setup now creates that context before anything is traced (tokamax._src.ops.op.get_autotuning_cache_overlay_state, a private hook, skipped when absent). Measured on an RTX Pro 4500 Blackwell MIG slice, 3 identical predictions in one process, no persistent cache: 71.7 / 67.7 / 18.2 s before, 71.8 / 18.2 / 18.2 s after, jit cache size 1 instead of 2, identical outputs. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
AlphaPulldown compiles its models from scratch in every process: on an H100 a 164-token AF3 prediction takes 51 s, 47 s of it compiling. AF3 also compiled the model a second time on the second prediction of every process. JAX's persistent compilation cache removes the per-process cost, but
--jax_compilation_cache_dirsegfaulted on shared filesystems. This PR makes the cache safe, turns it on by default, and removes the second compile.alphapulldown/prediction/jax_compilation_cache.py). It keeps only JAX's own entries (jax_persistent_cache_enable_xla_caches=none).JAX_PERSISTENT_CACHE_ENABLE_XLA_CACHESin the environment is respected.$JAX_COMPILATION_CACHE_DIR, else$XDG_CACHE_HOME/alphapulldown/jax_compilation_cache(~/.cache/...).none,falseor""turns the cache off.setup(jax_compilation_cache_dir=None)keeps caching off.filelock, which the images do not ship.jax.make_user_context) the first time an op consults its autotuning cache, during the first trace. A new user context changes every jit cache key, so the second predict call in a process re-traced and recompiled the whole model (~50 s); only later calls hit.setupnow creates the context before anything is traced. It uses a private tokamax hook, and is skipped when the hook is absent.Companion: KosinskiLab/AlphaPulldownSnakemake#57 turns the cache on in the workflow. It works with today's 2.9.0 images on its own, because it also exports the environment variable.
Validation
Unit suite (serial): 665 passed, 10 skipped, including 21 new tests (main: 644 passed).
test_package_layoutlegacy-import failures also occur on unmodifiedmainwhen run serially; CI's--dist loadfiledoes not show them.AF3, cache on BeeGFS
/scratch. RTX Pro 4500 Blackwell MIG slice, 2.9.0 image running this branch's package.Default on, through
run_structure_prediction:$XDG_CACHE_HOME, and the next process hits it (74.2 s → 29.0 s).--jax_compilation_cache_dir=none: nothing is cached.Second-prediction recompile: 3 identical predictions in one resident batch, no persistent cache.
AF2 multimer, 5 models, same GPU type:
The 5 models share compiled executables through the cache. Ranking scores and top-model atoms are identical.
AF3 benchmark on 9 GPU types (RTX 3090, A40, L40S, A100, H100 PCIe, H100 SXM, H200, RTX Pro 6000 and 4500 Blackwell):
🤖 Generated with Claude Code