Shared JAX compile cache for AF2 and AF3 at every batch size - #57
Merged
Merged
Conversation
Every structure_inference job now gets --jax_compilation_cache_dir (<output_directory>/.jax_compilation_cache) for both JAX backends, per-fold and resident alike. Without it each process recompiles its models, and AF3 recompiles on every prediction call even inside one resident batch (~50 s of a ~75 s call at 256 tokens). AF2 has accepted the flag since AlphaPulldown 2.8.0; the default images are 2.9.0. The rule exports JAX_PERSISTENT_CACHE_ENABLE_XLA_CACHES=none (unless set), so JAX keeps only its own entries in the cache. XLA's per-fusion autotune cache, which JAX writes there by default, fails its atomic rename on BeeGFS and XLA then segfaults: that is what crashed AF3 batching, and the reason the cache was kept away from resident batches and from AF2. Users move the cache with --jax_compilation_cache_dir in structure_inference_arguments, or turn it off with false/null/"". AlphaLink (PyTorch) gets no cache. The separate resident argument set is gone: it now equals the per-fold one. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
6 tasks done
…re, not per call The cache needs no action, so its note moves out of the always-visible backend-flags callout and the batching section into one collapsed block. AF3 does not recompile 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>
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
Turn on one shared JAX compile cache for every structure-inference job, for AlphaFold2 and AlphaFold3, at every
batch_size. Without it every inference process compiles its models from scratch: minutes per AlphaFold2 fold, and about 50 s per AlphaFold3 token bucket. At 256 tokens that is 50 s of a 75 s prediction. AlphaPulldown 2.9.1 and older also compile the first bucket twice; that is fixed in KosinskiLab/AlphaPulldown#644.batch_inference_argsnow adds--jax_compilation_cache_dir=<output_directory>/.jax_compilation_cachefor both JAX backends, per-fold and resident alike.structure_inferencerule exportsJAX_PERSISTENT_CACHE_ENABLE_XLA_CACHES=none(unless already set), so JAX keeps only its own entries in the cache.--jax_compilation_cache_dirinstructure_inference_argumentsto move the cache, or tofalse/null/""to turn it off.config.yamlupdated.Validation
CI-equivalent environment (Python 3.12, Snakemake 9.27,
alphapulldown-input-parser0.5.x): 157 passed, 2 skipped, including the AF2/AF3 DAG dry-runs. The 2 skips need the optional Slurm plugin.Rendered
structure_inferencecommands, from dry-runs with--printshellcmdson theresident_batchfixtures:batch_size1 and 2, all carry the cache flag and the export.--allow_resumeappears only for AF2 batches.AF3 on GPU, cache on BeeGFS
/scratch, 2.9.0 image with the environment variable (RTX Pro 4500 Blackwell MIG slice):AF2 multimer, 5 models:
Outputs are bit-identical.
AF3 benchmark on 9 GPU types: a prediction in a fresh process with a warm cache is 1.6–2.0× faster on average (3–4.5× at 164 tokens), with identical ranking scores.
🤖 Generated with Claude Code