TurboPairFormer:优化三角注意力内核,加速蛋白质折叠模型训练
TurboPairFormer: Fast and Stable Protein Folding Model Training with an Optimized Triangle Attention Kernel
搞蛋白质结构预测的可以看看,这个 TurboPairFormer 把 AlphaFold3 式训练的三角注意力算子重写了一遍,比 Triton 后端快 1.73 倍,数值还更稳定。
三角注意力是 AlphaFold3 类生物分子模型的核心计算,token 数增加时开销呈三次方增长。TurboPairFormer 针对 NVIDIA Hopper GPU 重新实现了这一计算,采用 key-tile 并行反向算法,避免浮点原子操作和完整 score 梯度存储。其输出残差补偿方法将各梯度相对 FP64 参考的 RMSE 降低 28-47%,且在全部 600 个输入案例中五次重复调用梯度 bitwise 一致。集成进 OpenFold3 后,在 16 张 H100 上 crop size 768 时比 OpenFold3 的 Triton 后端快 1.73 倍,比 cuEquivariance 快 1.13 倍。
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\times$ over OpenFold3's Triton backend and $1.13\times$ over cuEquivariance at crop size 768.