cs.DCSep 28, 2026

Accelerator Choice Is Not Enough: AlphaFold2 Inference on Cloud TPUs

Authors: Lorenzo Pazienza, Ihab El Bani

Organizations: LUISS Guido Carli University Rome, Italy · Stanford University · Al Akhawayn University Ifrane, Morocco

Abstract

AlphaFold2 is written in JAX, so the same inference code compiles and runs unchanged on CPUs, GPUs and Google Cloud TPUs. That portability makes the accelerator look like the main decision a user has to make. We show that it is not. Running one AlphaFold2 inference workload across a Colab CPU runtime, an NVIDIA T4 GPU and a dedicated eight-chip Cloud TPU v5e slice, we find a large hardware advantage for the TPU, 0.47 s per call in steady state on a single chip against 13.1 s on the T4 in the same measurement campaign, and three ways in which the software layer decides how much of it a user actually gets. The default execution path uses one chip of the eight, and at list prices the idle capacity makes the slice cost about as much per prediction as the GPU. Batching with JAX's vmap never exceeds single-query throughput, while mapping queries across chips with JAX's pmap gives eight chips 6.5-7.9x the throughput of one on a matched grid; automatic sharding leaves the per-chip footprint unchanged, consistent with replication, most plausibly because AlphaFold2 carries no sharding annotations. Our retained trace analysis of a first call at a new input shape reports about three quarters of the traced span in JAX tracing and compilation rather than execution. Reruns five weeks later reproduced neither cloud baseline, the GPU one off by roughly a factor of two, so the hardware ratio above is specific to one campaign.

Figures & tables

Explore similar work

Oct 5, 2026cs.LG

TurboPairFormer: Fast and Stable Protein Folding Model Training with an Optimized Triangle Attention Kernel

Triangular attention is a core computation in AlphaFold3-style biomolecular models, with cubic cost in token count. Its shared pair bias adds a gradient reduction across attention slices to the usual reductions over queries and keys. The open-source backends we examine handle these reductions through repeated probability recomputation, floating-point atomics, or full score-gradient storage. Separately, computing the softmax backward correction from BF16-rounded forward outputs loses numerical precision. We present TurboPairFormer, a triangular attention implementation for NVIDIA Hopper GPUs that addresses both issues. Our key-tile-parallel backward algorithm recomputes each probability tile once for the query, key, value, and pair-bias gradients, using ordered partial reductions for deterministic accumulation without floating-point atomics or full score-gradient storage. Output-residual compensation retains a BF16 approximation of the output-rounding residual to compute the backward correction more accurately in FP32, without changing the BF16 output. With BF16 inputs at crop sizes 384, 640, and 768 and head dimensions 16 and 32, TurboPairFormer achieves the lowest mean query, key, and pair-bias gradient RMSE against an FP64 reference among the implementations compared in this paper. Residual compensation reduces these RMSE values by 28-47% in controlled ablations. All four gradients are bitwise identical across five repeated calls in all 600 input cases under fixed execution conditions. Integrated into OpenFold3 with our triangle multiplication kernels, TurboPairFormer achieves the lowest GPU computation time per optimizer step among the evaluated backend configurations on 16 H100 GPUs, with speedups of 1.73×1.73\times over OpenFold3's Triton backend and 1.13×1.13\times over cuEquivariance at crop size 768.
Apr 28, 2026cs.LG

Block-Wise Differentiable Sinkhorn Attention: Tail-Refinement Gradients with a Gap-Aware Dustbin Bridge

We study long-context balanced entropic optimal transport (OT) attention on TPU hardware through a stopped-base, fixed-depth tail-refinement surrogate. After a stopped TT-step Sinkhorn solve, we unroll a short refinement tail and differentiate that surrogate exactly. For the reported R=2R=2 TPU path, the backward pass contains four staircase plan factors. We prove an exact one-reference-tile schedule: the R=2R=2 score cotangent is a single reference plan tile times an explicit modifier field built from vector cotangents and dual differences. This yields block-wise cost O((T+R)LW)O((T+R)LW), O(Ld)O(Ld) input storage, and O(L)O(L) additional HBM usage for fixed head dimension dd and band width WW on the balanced fixed-support path. We also formalize the current \texttt{dustbin_block} path as the same unit-target surrogate on an augmented support, so the adjoint schedule lifts to the single-active-dustbin path used in our TPU runs; this bridge is algebraic and does not claim a general KL-unbalanced or arbitrary-capacity gap model. We provide a local surrogate-bias bound, an a posteriori bias certificate, and a projective contraction certificate for strictly positive active blocks. On synthetic masked problems, the optimized kernel matches exact autodiff of the same centered surrogate to within 10−510^{-5}--10−1010^{-10}. On TPU v6e-8, a four-configuration Pfam screen completes end-to-end, and a promoted balanced R=2R=2 run sustains roughly 8.58.5 examples per second through a three-hour budget, reaching step 14371437. Held-out Pfam test shards improve reconstruction from 5.575.57 to 2.052.05 and sparse CE from 5.535.53 to 5.305.30 relative to step 00, with CE logged diagnostically rather than optimized directly; target-barycenter alignment metrics do not materially improve, and a deterministic diagonal reference remains stronger on those metrics.
Jul 24, 2026cs.AR

FusionML: Prefill, Not Decode - Mechanism and Boundaries of CPU+GPU Co-Execution on Unified-Memory Apple Silicon

Apple-Silicon SoCs share CPU, GPU, and Neural Engine over one unified memory system, raising the question of whether transformer inference can be accelerated by splitting single operators across units. Prior attempts, including our own, failed or produced precision-confounded wins. We identify the cause: MLX's lazy-graph scheduler \emph{serializes} cross-stream work whenever a CPU-stream operation consumes an unmaterialized GPU result inside one evaluation graph, so a row-split matmul that runs \x{1.38} faster with materialized inputs runs \x{0.66} slower than GPU-only inside a lazy graph; an eager materialization boundary restores concurrency (\x{1.34}). \sys{} implements a per-layer, contention-aware CPU+GPU row split for transformer prefill built on this fix. Evaluated across five chips and three Apple-Silicon generations, community-replicated, the split accelerates Llama-shaped decoder-block prefill by \x{1.15}--\x{1.38}, unchanged at full 32-block depth, and reaches \x{1.18}--\x{1.25} faster time-to-first-token on a real Qwen2.5-7B checkpoint served through stock MLX-LM, with token-identical outputs and unchanged decode throughput. We characterize the boundaries equally carefully: decode cannot benefit, bound by shared bandwidth co-execution does not add; precision-matched training loses \x{0.86}--\x{0.97} on all five chips; ANE dispatch overhead excludes it at layer granularity; and a no-regression runtime gate becomes self-defeating under memory pressure, where probing an alternative mode evicts the active mode's working set. Code, raw results, and generation transcripts are released.