TurboPairFormer: Fast and Stable Protein Folding Model Training with an Optimized Triangle Attention Kernel
Organizations: Stevens Institute of Technology · Lambda · Open Molecular Software Foundation
Abstract
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 over OpenFold3's Triton backend and over cuEquivariance at crop size 768.
Figures & tables
Appendix figures & tables20 assets
Supplementary material from the paper’s appendix.
Appendix
| Notation | Meaning |
|---|---|
| , , | Examples per GPU, tokens per crop, and number of fixed-token attention slices. The label in denotes the outer axis; . |
| , , | Heads, features per head, and features per pair; the label refers to the pair tensor . |
| ; | Token indices: the fixed token (outer index), and the query and key positions within its slice. The triangle example uses and . A subscript alone selects one fixed-token slice. |
| ; | Feature coordinate; key tile, outer group, and forward 16-key chunk indices. There are four chunks per forward tile. In the memory analysis, denotes an execution point. |
| Pair features of shape . selects one pair, and (also written ) the slice of a fixed token ; numeric subscripts such as select individual pairs. | |
| Query, key, value, and attention output tensors, each of shape ; for one example, slice, and head, each is an matrix. |
| 384 | 216 | 108 | 108 | 36 | 864 |
| 640 | 1000 | 500 | 500 | 100 | 4000 |
| 768 | 1728 | 864 | 252 | 144 | 6912 |
| 1024 | 4096 | 2048 | 256 | 256 | 16384 |
| Backend | Auto- | No dense | No | No bias | Fixed | |
|---|---|---|---|---|---|---|
| tune | recomputations | FP32 | atomics | atomics | order | |
| TurboPairFormer | No | 1 | Yes | Yes | Yes | Yes |
| TriFast | Yes | 3 | Yes | Yes | Yes | Yes |
| Protenix | Yes | 3 | Yes | Yes | Yes | Yes |
| MegaFold | No | 2 | Yes | Yes | No | No |
| OpenFold3 Triton (newer) | No | 2 | Yes | No | Yes | No |
| Backend | Crop 384 | Crop 640 | Crop 768 |
|---|---|---|---|
| TurboPairFormer | 7.01 | 11.17 | 14.67 |
| Protenix | 7.35 | 11.30 | 14.92 |
| cuEquivariance | 7.43 | 12.11 | 16.51 |
| TriFast | 7.51 | 11.44 | 15.14 |
| MegaFold | 7.63 | 13.11 | 18.02 |
| DeepSpeed-Evo | 7.52 | 13.34 | 18.42 |
| Setting | Post-warmup study | Fine-tuning validation |
| Starting weights | Preview2 after 1,000 updates at crop 640 | Preview2 |
| GPUs | One H100 | 16 H100 (two nodes) |
| Crop size | 384, 640, 768 | 384 |
| Training updates | 400 | 1,000 |
| Learning rate | , post-warmup plateau | , 50 warmup steps |
| EMA decay | 0.999 | 0.98 |
| Quantity | Protenix | cuEquivariance | TurboPairFormer |
|---|---|---|---|
| 0.1695 | 0.1686 | 0.1686 | |
| 0.8202 | 0.8222 | 0.3769 | |
| 0.4347 | 0.4358 | 0.2437 | |
| 0.2201 | 0.2164 | 0.2164 | |
| 0.3534 | 0.3506 | 0.2449 |
| Checkpoint | Difference | 95% interval |
|---|---|---|
| Preview2 (step 0) | ||
| PyTorch | ||
| OpenFold3 Triton | ||
| cuEquivariance | ||
| Protenix | ||
| TriFast |
| Forward | Backward | Forward + backward | Change vs. cuEq (%) | ||||||||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| Ours | cuEq | Protenix | Ours | cuEq | Protenix | Ours | cuEq | Protenix | Fwd | Bwd | F+B | ||
| 16 | 384 | 0.384 | 0.640 | 0.417 | 1.167 | 2.097 | 1.285 | 1.493 | 2.623 | 1.587 | |||
| 16 | 640 | 1.177 | 2.041 | 1.316 | 4.300 | 8.046 | 4.655 | 5.398 | 9.918 | 5.830 | |||
| 16 | 768 | 1.870 | 3.341 | 2.137 | 7.392 | 13.244 | 7.563 | 9.025 | 16.478 | 9.499 | |||
| 32 | 384 | 0.474 | 0.708 | 0.502 | 1.520 | 2.430 | 1.482 | 1.920 | 3.015 | 1.854 | |||
| 32 | 640 | 1.538 | 2.192 | 1.602 | 5.468 | 8.764 | 5.654 | 6.810 | 10.823 | 7.129 | |||
| Forward time (ms) | Forward reduction (%) | F+B reduction (%) | |||||
|---|---|---|---|---|---|---|---|
| Baseline | Core | Combined | Core | Combined | Core | Combined | |
| 384 | 0.547 | 0.408 | 0.388 | 25.5 | 29.1 | 8.1 | 9.2 |
| 640 | 2.037 | 1.478 | 1.426 | 27.4 | 30.0 | 8.3 | 9.0 |
| 768 | 3.336 | 2.419 | 2.341 | 27.5 | 29.8 | 7.8 | 8.5 |
| Incoming | Outgoing | ||||
|---|---|---|---|---|---|
| Channels | cuEquivariance | Ours | cuEquivariance | Ours | |
| 384 | 128 | 2.769 | 1.326 ( ) | 2.777 | 1.323 ( ) |
| 640 | 128 | 3.615 | 3.334 ( ) | 3.678 | 3.336 ( ) |
| 768 | 128 | 5.142 | 4.788 ( ) | 5.210 | 4.782 ( ) |
| 384 | 64 | 2.734 | 1.111 ( ) | 2.735 | 1.116 ( ) |
| 640 | 64 | 2.777 | 1.921 ( ) | 2.784 | 1.922 ( ) |
| Thread block owns | Local | Cross-block | Partials | Thread blocks |
|---|---|---|---|---|
| outer indices per thread block | ||||
| Key tile (ours) | , | , | ||
| Query tile | , , | |||
| Whole slices (DeepSpeed-Evo) | , , | none | ||
| All outer indices per thread block | ||||
| Query–key tile | , , | |||
| Change with non-deterministic atomic accumulation (%) | |||||||
| Deterministic (ms) | Atomic | Per-element | Both, per-element | TMA bulk | Both, TMA bulk | ||
| 16 | 384 | 0.849 | |||||
| 16 | 640 | 3.582 | |||||
| 16 | 768 | 6.349 | |||||
| 32 | 384 | 1.091 | |||||
| 32 | 640 | 4.505 | |||||
| Crop | Attention only | TriMul only | Both |
|---|---|---|---|
| 384 | 2.21 | 2.86 | 5.55 |
| 640 | 5.59 | 2.78 | 8.45 |
| 768 | 6.43 | 2.99 | 9.63 |