Skip to content

Shared JAX compile cache for AF2 and AF3 at every batch size - #57

Merged
DimaMolod merged 4 commits into
mainfrom
exp/compile-cache
Oct 5, 2026
Merged

DimaMolod merged 4 commits into
mainfrom
exp/compile-cache

Conversation

@DimaMolod

@DimaMolod DimaMolod commented Oct 5, 2026 •

Copy link
Copy Markdown
Collaborator

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.

  • Where the cache goes: batch_inference_args now adds --jax_compilation_cache_dir=<output_directory>/.jax_compilation_cache for both JAX backends, per-fold and resident alike.
    • AlphaFold2 has accepted the flag since AlphaPulldown 2.8.0; the default images are 2.9.0.
    • AlphaLink (PyTorch) gets no cache.
  • What makes BeeGFS safe: the structure_inference rule exports JAX_PERSISTENT_CACHE_ENABLE_XLA_CACHES=none (unless already set), so JAX keeps only its own entries in the cache.
  • User control: set --jax_compilation_cache_dir in structure_inference_arguments to move the cache, or to false/null/"" to turn it off.
  • Simplification: the separate resident-batch argument set is gone, because it now equals the per-fold one.
  • Docs: the README cache note is one collapsed block (no action needed), matching AlphaPulldown's; flag reference and config.yaml updated.

Validation

  • CI-equivalent environment (Python 3.12, Snakemake 9.27, alphapulldown-input-parser 0.5.x): 157 passed, 2 skipped, including the AF2/AF3 DAG dry-runs. The 2 skips need the optional Slurm plugin.

  • Rendered structure_inference commands, from dry-runs with --printshellcmds on the resident_batch fixtures:

    • AF2 and AF3, at batch_size 1 and 2, all carry the cache flag and the export.
    • --allow_resume appears 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):

    • A second process with a different complex in the same token bucket hits the cache: 75 s → 30 s per call.
    • Without the variable, the same setup segfaults.
  • AF2 multimer, 5 models:

    fold no cache cache
    first fold 480 s 255 s
    same shape again 480 s 152 s
    new length 621 s 419 s

    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

DimaMolod and others added 2 commits October 5, 2026 11:33
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>
DimaMolod and others added 2 commits October 5, 2026 13:36
…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>
@DimaMolod
DimaMolod merged commit 0a594b8 into main Oct 5, 2026
2 checks passed
@DimaMolod
DimaMolod deleted the exp/compile-cache branch October 5, 2026 12:34
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant