AlphaFold2 runs faster on cloud TPUs but software limits speed gains
Accelerator Choice Is Not Enough: AlphaFold2 Inference on Cloud TPUs
Distributed, Parallel, and Cluster ComputingMachine LearningPerformance
Summary
The authors find that using Google Cloud TPUs can make AlphaFold2 run much faster than on GPUs or CPUs. However, the software setup often means users only get the speed of one TPU chip, not the full power of many chips available. Also, most of the time is spent translating AlphaFold2’s code into machine instructions rather than actual running. This shows that just picking the right hardware is not enough to get the best speed from AlphaFold2.
What this means in practice
- •For cloud platform engineers: Optimize AlphaFold2 deployment by adjusting software parallelism to better utilize multiple TPU chips and improve throughput.
- •For bioinformatics infrastructure teams: Choose hardware and configure AlphaFold2 setups informed by detailed performance trade-offs between CPUs, GPUs, and TPUs for protein folding tasks.
Authors
Lorenzo Pazienza, Ihab El Bani
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.vmap never exceeds single-query throughput, while mapping queries across chips with jax.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.