Skip to content

Safe persistent JAX compile cache, on by default; no AF3 recompile on the second prediction - #644

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

DimaMolod merged 5 commits into
mainfrom
exp/compile-cache

Conversation

@DimaMolod

@DimaMolod DimaMolod commented Oct 5, 2026 •

Copy link
Copy Markdown
Collaborator

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_dir segfaulted on shared filesystems. This PR makes the cache safe, turns it on by default, and removes the second compile.

  • One helper for both backends (alphapulldown/prediction/jax_compilation_cache.py). It keeps only JAX's own entries (jax_persistent_cache_enable_xla_caches=none).
    • Why: by default JAX also writes XLA's per-fusion autotune cache into the directory. Inserting into it fails an atomic rename on BeeGFS ("Failed to insert autotune cache: FAILED_PRECONDITION") and XLA then segfaults. This is what crashes AF3 batching in AlphaPulldownSnakemake.
    • AF2: jax 0.5.3 carries the same per-fusion cache, so AF2 gets the same fix.
    • An explicit JAX_PERSISTENT_CACHE_ENABLE_XLA_CACHES in the environment is respected.
  • No crash on a bad directory: if the cache directory cannot be created or written, the helper logs a warning and runs without a cache.
  • On by default for command-line runs. An unset flag resolves to $JAX_COMPILATION_CACHE_DIR, else $XDG_CACHE_HOME/alphapulldown/jax_compilation_cache (~/.cache/...).
    • none, false or "" turns the cache off.
    • Programmatic setup(jax_compilation_cache_dir=None) keeps caching off.
  • The per-user default is capped at 10 GB, pruned oldest-first when a backend starts.
    • Why: AF3 adds about 3 MB per token bucket, but AF2 adds about 5 MB per new complex size, without bound.
    • Safety: pruning is lock-free and best effort; an entry that disappears only costs a recompile.
    • Scope: a directory the user chose is never pruned. jax's own size cap needs filelock, which the images do not ship.
  • AF3 no longer recompiles on its second prediction.
    • Cause: tokamax creates a JAX user context (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.
    • Fix: setup now creates the context before anything is traced. It uses a private tokamax hook, and is skipped when the hook is absent.
    • Vanilla AlphaFold 3 likely has the same issue (untested).
  • README: the cache note is now one collapsed block, matching AlphaPulldownSnakemake's.

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).

    • The 2 test_package_layout legacy-import failures also occur on unmodified main when run serially; CI's --dist loadfile does not show them.
  • AF3, cache on BeeGFS /scratch. RTX Pro 4500 Blackwell MIG slice, 2.9.0 image running this branch's package.

    • 2.9.0 as released segfaults (rc=139).
    • This branch runs cleanly.
    • A different complex in the same token bucket, in a new process, gets a cache hit: 76.6 s → 29.2 s per call.
  • Default on, through run_structure_prediction:

    • No flag: the cache is created under $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.

    • Before: 71.7 / 67.7 / 18.2 s.
    • After: 71.8 / 18.2 / 18.2 s.
    • jit cache size 1 instead of 2; identical ranking scores.
  • AF2 multimer, 5 models, same GPU type:

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

    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):

    • A prediction in a fresh process with a warm cache is 1.6–2.0× faster on average: 3–4.5× at 164 tokens, about 1.1× above 2,500 tokens.
    • All 39 cached predictions rank their samples exactly as the compiling run did.

🤖 Generated with Claude Code

DimaMolod and others added 2 commits October 5, 2026 10:13
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>
@chatgpt-codex-connector

Copy link
Copy Markdown

You have reached your Codex usage limits for code reviews. You can see your limits in the Codex usage dashboard.

DimaMolod and others added 3 commits October 5, 2026 13:36
…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>
@DimaMolod DimaMolod changed the title Safe persistent JAX compile cache, on by default Safe persistent JAX compile cache, on by default; no AF3 recompile on the second prediction Oct 5, 2026
@DimaMolod
DimaMolod merged commit 9ed1b89 into main Oct 5, 2026
6 checks passed
@DimaMolod
DimaMolod deleted the exp/compile-cache branch October 5, 2026 12:40
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