前向卡在 exp,后向卡在共享内存流量——FA-4 如何用 TMEM、2-CTA MMA、DSMEM 和确定性模式把后向的 smem 瓶颈压下去
第 6 课讲了前向:非对称缩放让 exp 成瓶颈,FA-4 用软仿真 exp 和条件 rescale 应对。本课讲后向——瓶颈换成了共享内存流量,解法也完全不同 [FA-4 论文 §3.2]。最后覆盖架构无关的 LPT 调度和实现语言 CuTe-DSL,为 FA-1/2/3/4 系列收官。
前向只有 2 个 GEMM(QKᵀ 和 PV),后向要算 5 个 MMA:重算 S,再算 QK 的梯度 dQ/dK、PV 的梯度 dP/dV,外加 dsoftmax 的 elementwise [论文 §3.2.1]。算力是前向的约 2.5×。
但 Blackwell 的非对称缩放在这里换了个方向发难。后向 roofline(M=N=d=128)[论文 §3.2.1, Table 3]:
| 资源(M=N=d=128,1-CTA) | 周期数 |
|---|---|
| MMA 计算(5 个 GEMM) | 2560 |
| exp 单元 | 1024 |
| 共享内存读(MMA 操作数) | 2048 |
| 共享内存(dS 写 + dQ 读写) | 1280 |
| 共享内存总计 | 3328 |
前向(第 6 课):exp 1024 = MMA 1024,exp 是瓶颈。
后向:smem 3328 > MMA 2560,共享内存流量超 MMA 30%。8 个 BF16 操作数要从 smem 读进 Tensor Core,这个搬运比算还慢 [Together AI 博客]。
所以后向的优化目标不是抬算力,而是砍 smem 流量,并让剩下的非 matmul 活和 MMA 重叠。
FA-3 后向之所以序列化,是因为累加器放在寄存器里——寄存器是稀缺资源,同时放不下 5 套累加器,只能按 S→dP→dV→dQ→dK 的顺序一个一个来,只有 TMA load 能显著乱序 [论文 §3.2.2]。
FA-4 把累加器挪进 TMEM(256 KB/SM,见第 6 课),寄存器压力解除,多个 MMA 能同时 in flight。关键重叠是 [论文 §3.2.2] [Together AI 博客]:
# 概念伪代码:后向主循环,沿 KV 维度迭代
for j in range(T_kv):
# —— tile j 的 softmax ——
recompute S_j, P_j # 2 个 MMA
dP_j = dsoftmax(...) # elementwise,是瓶颈段
# —— 与此同时,发上一轮 tile j-1 的 dK、dQ ——
dK[j-1] = dS[j-1]^T · Q # MMA(累加器在 TMEM,不抢寄存器)
dQ[j-1] = dS[j-1] · K # MMA(atomic 累加)
# ↑ softmax[j] 和 MMA[j-1] 并行,把 softmax 藏起来
一个 SM 的 TMEM 最多放 4 个 128×128 累加器,但后向有 5 个累加器(S、P、dP、dS、dQ),其中 dV、dK 要跨迭代累加不能共享。FA-4 的解法是分时复用:S 和 P 共用一组 TMEM 列(offset 0),dP、dS、dQ 共用另一组 [论文 §3.2.2]。
后向还把 S、P 重算成转置 tile(Sᵀ、Pᵀ),这样它们在 TMEM 里恰好就是 dV、dK 这两个 TS-MMA 需要的 operand-A 布局,省掉额外搬运 [论文 §3.2.2]。
即便两个操作数搬进 TMEM,后向 5 个 GEMM 还有 8 个 BF16 操作数要从 smem 读,smem 流量仍超 MMA 30%。FA-4 用 Blackwell 的 2-CTA MMA 模式砍这部分流量 [论文 §2.2, §3.2.3]。
2-CTA 模式下,一个 cluster 里的两个 CTA 合作完成一次 MMA:把输出累加器在 M 维劈成两半,每个 CTA 持一半;关键是 operand B 也劈成两半,每个 CTA 只 stage 自己那一半 B [论文 §2.2] [Together AI 博客]。
| 配置 | 1-CTA | 2-CTA |
|---|---|---|
| MMA tile M | 128 | 256 |
| 每 CTA 持有 B | 全部 | 一半 |
| smem 流量(MMA 操作数) | 2048 cyc | 1536 cyc |
| smem 总计 | 3328 cyc | 2688 cyc |
2-CTA 把 smem 总流量从 3328 压到 2688,和 MMA 的 2560 基本拉平,瓶颈被消掉 [论文 §3.2.1, Table 3]。代价是 kernel 里 TMEM 和 Tensor Core 操作必须全程保持 2-CTA 模式,CTA 要成对启动 [论文 §2.2]。
2-CTA 带来一个 dQ 的归约轴冲突 [论文 §3.2.3] [Together AI 博客]:
解法是用 cluster 内的 DSMEM(Distributed Shared Memory)交换半块 dS [论文 §3.2.3]:
(M/2, 2N) × (2N, d),归约维 2N 在单个 CTA 内完整——能用 2-CTA 模式了。附带收益:dQ 的 global atomic 减半。
dQ 是跨迭代累加的,每轮写回靠 global atomic add(非确定、昂贵)。2-CTA 下每个 CTA 只写 M/2 行 → atomic 次数减半 [论文 §3.2.3]。
为了藏掉 DSMEM 交换的延迟,FA-4 重排流水线:先算当前 tile 的 dP,再算上一 tile 的 dQ。这样 dS 的 elementwise 和上一轮的 dQ MMA 重叠,DSMEM 延迟被盖住 [论文 §3.2.3]。
后向的非确定性根源就是上面那个 global atomic dQ 累加——多个 CTA 的累加顺序不定。对可复现训练(尤其强化学习)这是问题 [论文 §3.2.4]。
FA-4 提供确定性模式 [论文 §3.2.4] [Together AI 博客]:
强化学习的 reward 信号噪声大,如果 attention backward 本身有数值抖动,很难区分"是策略在变"还是"kernel 不确定"。确定性模式把这一层噪声抹掉,是 FA-4 对 RL 场景的针对性支持。
前面所有技术都依赖 Blackwell 专属硬件。但 FA-4 有一个完全架构无关的优化——LPT 调度,FA-3 里已有,FA-4 强化 [论文 §3.3]。
问题:causal mask 和变长序列让不同 worktile 的主循环长度不同,负载不均,尾部拖慢整体 [论文 §3.3]。
LPT 调度是 FA-4 里唯一不依赖 Sm100 的技术,纯软件调度策略,理论上任何平台都能用(FA-3 已采纳)。其他 TMEM/2-CTA/DSMEM/确定性锁都是 Blackwell 专属。这条边界和第 5 课的 PPU810E caveat 一致。
FA-4 整个用 CuTe-DSL 实现——CUTLASS 的 Python kernel DSL [论文 §4] [Together AI 博客]:
这不是性能优化,而是工程生产力改进——让研究者不用精通 C++ 模板元编程也能改 attention kernel。安装:pip install flash-attn-4,用法 from flash_attn.cute import flash_attn_func [官方仓库]。
七节课下来,一条主线贯穿 1–4 代——每代解决上一代的瓶颈,结果瓶颈换个地方再出现:
| 代 | 目标硬件 | 解决的瓶颈 | 新瓶颈 |
|---|---|---|---|
| FA-1 | A100 | HBM 读写量(S/P 矩阵落地) | 并行度、非 matmul 开销 |
| FA-2 | A100 | loop order、并行度、非 matmul | Tensor Core 空闲、softmax bubble |
| FA-3 | H100 | Tensor Core 利用率(WGMMA/TMA/异步重叠/FP8) | exp 单元、共享内存带宽未随 MMA 翻倍 |
| FA-4 | B200 | 非对称缩放(软仿真 exp / TMEM / 2-CTA / LPT) | ——(截至 2026) |
更完整的对比(年份、论文、精度、性能数字)见 1-4 代演进速查,术语查 术语表。
到这里,FlashAttention 1–4 代的原理与演进主线已完整覆盖。MISSION 的核心目标——"逐代比较优化重点、目标硬件和瓶颈转移"——达成。后续若感兴趣,可深入 backward recomputation 的 CUDA 实现细节,或在 PPU810E 上做实测选型专题(见 NOTES.md 待定项)。
本课推荐阅读 Together AI 官方博客:FlashAttention-4(含后向 pipeline 图、2-CTA/DSMEM 示意图),再读原论文 FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling 的 §3.2(后向)、§3.3(调度)、§4(CuTe-DSL)。代码在 flash_attn/cute。
FA-4 的 TMEM、2-CTA、DSMEM、确定性锁都是 Sm100(Blackwell)专属,PPU810E(Sm80)不可用(见 第 5 课)。本课的 B200 数字对你都是理论值。LPT 调度是唯一可移植的。若要在 PPU810E 上实测选型,参考第 5 课的 FA-3 经验:CUTLASS 通用路径 + torch.compile 陷阱。
2-CTA 的 M/N 分裂细节、DSMEM 交换的时序、确定性模式的锁粒度——任何不清楚的地方随时问我。FlashAttention 主线到此完结,但 backward CUDA 细节和 PPU810E 实测专题仍可作为后续方向。