Skip to content

--fast_kernels: ColabFold's fused kernels for AF2-Multimer (~2x); fix padded AF2 scores - #645

Merged
DimaMolod merged 6 commits into
mainfrom
exp/af2-fused-kernels
Oct 6, 2026
Merged

DimaMolod merged 6 commits into
mainfrom
exp/af2-fused-kernels

Conversation

@DimaMolod

@DimaMolod DimaMolod commented Oct 6, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

--fast_kernels off|on|auto (AlphaFold 2 only, default off): about 2× faster AlphaFold-Multimer inference on NVIDIA GPUs of compute capability 8.0+. It uses ColabFold's maintained fused Pallas kernels (colabfold-kernels, MIT) through the new hooks in the alphafold fork (KosinskiLab/alphafold#13).

  • Checked before any model is built (alphapulldown/prediction/fast_kernels.py). AlphaPulldown runs on many clusters, so the check covers:
    • the package is installed;
    • the fork has the hooks;
    • the GPU is NVIDIA with compute capability ≥ 8.0;
    • one small kernel actually compiles and runs on this GPU.
    • on raises with the missing requirement; auto falls back to the standard code with a log line.
  • Multimer models only. Monomer configs run fp32 and the kernels are bf16, so monomers keep the standard code.
  • Traceable: each prediction directory gets inference_kernels.json (fused_kernels, compute_capability), read back from each runner's own config, so fast and stock results can be told apart in ranking and analysis.
  • Install:
    • pip install "alphapulldown[fast-kernels]"; the AF2 image installs the package.
    • Its small V100/T4 binary dependency needs glibc ≥ 2.28; --no-deps works where that's a problem.
  • YAML-safe: true/false are read as on/off. PyYAML turns a bare on in the Snakemake config into a boolean.
  • Submodule: alphafold is at the merged fork, ad317b4 (Optional fused Pallas kernels from colabfold-kernels (off by default) alphafold#13). Its tree is identical to the tested a20e833.

Bug fix, separate commit 4519a56a: padded AF2 predictions reported scores over the padding.

  • Symptom: with --desired_num_res, pLDDT, PAE, pTM, ipTM and the ranking confidence were computed over the padded length. A 599-residue dimer padded to 700 reported ipTM = pTM = 0.9993 (true 0.927 / 0.940) and mean pLDDT 86.9 (true 95.7), and kept a 700×700 PAE. The structures were fine.
  • Who is affected: AlphaPulldownSnakemake pads every AF2-multimer batch to its largest fold when batch_size > 1 (since 2.8.0), so every smaller fold in such a batch got these scores.
  • Cause: recalculate_confidence only trimmed when the PAE head's raw logits were present. RunModel replaces them with the computed PAE array, so it returned early every time.
  • Fix: pTM and ipTM are recomputed from aligned_confidence_probs over the real residues, using AlphaFold's own predicted_tm_score. Chain IDs come from the unpadded features; pLDDT, PAE and the probabilities are trimmed. Unpadded predictions are unchanged.

Validation

  • Unit suite (serial): 698 passed, 10 skipped, including new tests for the flag, the requirement checks and the padded scores. One of them uses AlphaFold's real confidence.py to show that recomputed padded scores equal the unpadded truth to 1e-6.

    • The 2 test_package_layout failures also occur on unmodified main when run serially.
  • Speed vs AlphaPulldown stock (model_1_multimer_v3, 3 recycles, compile excluded, geometric mean over 164–2,546 tokens):

    RTX 3090 A40 L40S A100 H100 PCIe H100 SXM RTX Pro 6000 RTX Pro 4500 MIG
    2.09× 2.66× 2.22× 2.37× 2.26× 2.42× 2.18× 1.95×

    That's 1.0–1.18× faster than ColabFold 1.6.3's own --use-fast-kernels, which is 1.8–2.4× on the same inputs. H200 is queued; I'll add it.

  • Accuracy: 12 heterodimers released after AF2-Multimer's training cutoff, DockQ against experimental structures, on H100, A40 and Blackwell MIG; stock seeds 0+1, kernels seeds 0 (+1 on H100/A40).

    • Confident complexes (> 0.7): every arm agrees within 0.01 ranking confidence.
    • Low-confidence complexes sometimes flip between modes. Stock code does the same: 8B2R's confident-but-wrong mode, which kernels hit once (0.63), also comes from ColabFold stock (0.70) on the same card.
    • 8SM0: a 10-seed sweep on A40 and H100 finds its confident mode in 1 of 20 kernel runs and 0 of 20 stock runs.
  • Kernels off is bit-identical to main under deterministic XLA (A40).

  • --fast_kernels=auto through run_structure_prediction:

    • enables the kernels on 3090 (8.6), L40S (8.9), H100 (9.0) and Blackwell MIG (12.0), all with this image's jax 0.5.3;
    • a 599-token dimer scores 0.929–0.930 on all four;
    • on with a monomer keeps all 5 monomer models on the standard code.
  • Padded scores, A40, deterministic XLA:

    • single fold: padded 599→700 now reports ipTM 0.9268 / pTM 0.9403 / pLDDT 95.65 over 599 residues; unpadded is 0.9269 / 0.9404 / 95.65.
    • workflow batch path (run_structure_prediction_batch, --desired_num_res=599, 357- and 599-token folds): each fold's scores are over its own residues.
    • The padded small fold still differs from its unpadded run by about 0.01 ipTM, likely because AF2's in-model MSA sampling draws shape-dependent random numbers. That is seed-like and pre-existing.

Companion: KosinskiLab/AlphaPulldownSnakemake#58 (the workflow accepts the flag).

🤖 Generated with Claude Code

DimaMolod and others added 4 commits October 5, 2026 15:25
… kernels)

--fast_kernels off|on|auto (AF2 only, default off) switches on the fused
kernels of the colabfold-kernels package through the alphafold fork's new
hooks (alphafold.model.fused_kernels): flash attention, fused LayerNorm and
triangle multiplication's fused projection. About 2x faster AF2-Multimer on
NVIDIA GPUs of compute capability 8.0+.

alphapulldown/prediction/fast_kernels.py checks every requirement before a
model is built, so other clusters get a clear answer instead of a failure
mid-prediction: the package, the fork hooks, an NVIDIA GPU of cc >= 8.0, and
one small kernel actually compiled and run. "on" raises with the missing
requirement; "auto" falls back to the stock code with a log line.

Only multimer configs are touched (monomers run fp32, the kernels are bf16).
Each prediction directory gets inference_kernels.json, read back from the
runner's own config, so fast and stock results can be told apart.
pyproject gains the fast-kernels extra; the AF2 image installs the package.
The alphafold submodule points at the fork branch exp/af2-fused-kernels.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
PyYAML reads an unquoted on/off in the Snakemake config as a boolean, which
reaches the command line as True/False. The flag is now a string normalised
by fast_kernels.normalise_mode; anything else still fails at backend setup.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
…mers keep stock code

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
With --desired_num_res (and so in every AlphaPulldownSnakemake AF2-multimer
batch with batch_size > 1, which pads each fold to the batch's largest),
pLDDT, PAE, pTM, ipTM and the ranking confidence were computed over the
padded length. A 599-residue dimer padded to 700 reported ipTM = pTM =
0.9993 instead of 0.927 / 0.940, mean pLDDT 86.9 instead of 95.7, and kept
a 700x700 PAE in its pickle and JSON. The structures were unaffected.

recalculate_confidence only trimmed when the PAE head's raw logits were
present; AlphaFold's RunModel replaces them with the computed PAE array, so
it returned early every time. It now also handles that case: pTM and ipTM
are recomputed from aligned_confidence_probs over the real residues with
AlphaFold's own predicted_tm_score (softmax(log p) = p, bin edges rebuilt
from the bin count and max_predicted_aligned_error), chain IDs come from
the unpadded features, and pLDDT, PAE and the probabilities are trimmed.
Unpadded predictions are left unchanged.

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 2 commits October 6, 2026 09:07
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
ad317b4 merges a20e833, which this branch pointed at; the tree is unchanged.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
@DimaMolod
DimaMolod merged commit 62488d1 into main Oct 6, 2026
6 checks passed
@DimaMolod
DimaMolod deleted the exp/af2-fused-kernels branch October 6, 2026 07:19
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