arXivDaily arXiv每日学术速递 周一至周五更新
arXiv周末暂无论文更新,休息一下吧,周末愉快~~

加速器选择并非全部:AlphaFold2 在云 TPU 上的推理

Accelerator Choice Is Not Enough: AlphaFold2 Inference on Cloud TPUs

Lorenzo Pazienza, Ihab El Bani

arXiv 2609.34818首次发表:更新:

发表机构

LUISS Guido Carli University; Al Akhawayn University(路易吉·圭多·卡利大学(卢伊斯大学); 阿卡韦恩大学)

机构由 AI 辅助整理,请以论文原文为准。

AI 中文总结

本文通过对比实验表明,AlphaFold2 推理性能不仅取决于加速器选择,软件层(执行路径、批处理、分片)同样关键,并量化了 TPU 与 GPU 的吞吐差异及跟踪编译开销。

AI 中文摘要

AlphaFold2 使用 JAX 编写,因此相同的推理代码可以在 CPU、GPU 和 Google Cloud TPU 上无需修改地编译和运行。这种可移植性使得加速器看起来像是用户必须做出的主要决策。我们表明事实并非如此。在 Colab CPU 运行时、NVIDIA T4 GPU 和专用的八芯片 Cloud TPU v5e 切片上运行一个 AlphaFold2 推理工作负载,我们发现 TPU 具有巨大的硬件优势,在相同测量活动中,单芯片稳态下每次调用 0.47 秒,而 T4 为 13.1 秒,并且软件层以三种方式决定用户实际获得多少优势。默认执行路径使用八个芯片中的一个,按标价计算,空闲容量使得每个预测的切片成本与 GPU 相当。使用此 http URL 进行批处理从未超过单查询吞吐量,而使用此 http URL 将查询映射到芯片上,在匹配的网格上,八个芯片的吞吐量是单个芯片的 6.5-7.9 倍;自动分片保持每芯片占用不变,与复制一致,最可能的原因是 AlphaFold2 没有分片注释。我们对新输入形状的第一次调用进行的保留跟踪分析报告,约四分之三的跟踪跨度用于 JAX 跟踪和编译,而非执行。五周后重新运行未能重现任一云基线,GPU 基线偏差约两倍,因此上述硬件比率仅特定于一次活动。

英文摘要

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.

Comments22 pages, 5 figures, 4 tables. Code and data: https://github.com/lorenzopazienza/alphafold-tpu-benchmark

论文原文

arXiv 摘要页 · PDF 原文 · HTML 原文

↑