第 7 课:FlashAttention-4(下)后向共享内存瓶颈与 LPT 调度

前向卡在 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 系列收官。

1. 后向瓶颈转移:从 exp 到共享内存流量

前向只有 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
关键直觉:前向 vs 后向

前向(第 6 课):exp 1024 = MMA 1024,exp 是瓶颈。
后向:smem 3328 > MMA 2560,共享内存流量超 MMA 30%。8 个 BF16 操作数要从 smem 读进 Tensor Core,这个搬运比算还慢 [Together AI 博客]

所以后向的优化目标不是抬算力,而是砍 smem 流量,并让剩下的非 matmul 活和 MMA 重叠。

2. TMEM 累加器:让多个 MMA 同时 in flight

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 藏起来
TMEM 列复用:放不下 5 套累加器

一个 SM 的 TMEM 最多放 4 个 128×128 累加器,但后向有 5 个累加器(S、P、dP、dS、dQ),其中 dV、dK 要跨迭代累加不能共享。FA-4 的解法是分时复用SP 共用一组 TMEM 列(offset 0),dPdSdQ 共用另一组 [论文 §3.2.2]

后向还把 S、P 重算成转置 tile(Sᵀ、Pᵀ),这样它们在 TMEM 里恰好就是 dV、dK 这两个 TS-MMA 需要的 operand-A 布局,省掉额外搬运 [论文 §3.2.2]

3. 2-CTA MMA:砍掉一半 operand-B 流量

即便两个操作数搬进 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-CTA2-CTA
MMA tile M128256
每 CTA 持有 B全部一半
smem 流量(MMA 操作数)2048 cyc1536 cyc
smem 总计3328 cyc2688 cyc

2-CTA 把 smem 总流量从 3328 压到 2688,和 MMA 的 2560 基本拉平,瓶颈被消掉 [论文 §3.2.1, Table 3]。代价是 kernel 里 TMEM 和 Tensor Core 操作必须全程保持 2-CTA 模式,CTA 要成对启动 [论文 §2.2]

4. DSMEM 交换 + dQ atomic 减半

2-CTA 带来一个 dQ 的归约轴冲突 [论文 §3.2.3] [Together AI 博客]

解法是用 cluster 内的 DSMEM(Distributed Shared Memory)交换半块 dS [论文 §3.2.3]

  1. 两个 CTA 互换各自一半的 dS,重新 pack 成"每个 CTA 持 M/2 行 × 全 2N"的布局。
  2. 于是 dQ 的 MMA 变成 (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]

5. 确定性模式:可复现训练不丢吞吐

后向的非确定性根源就是上面那个 global atomic dQ 累加——多个 CTA 的累加顺序不定。对可复现训练(尤其强化学习)这是问题 [论文 §3.2.4]

FA-4 提供确定性模式 [论文 §3.2.4] [Together AI 博客]

为什么 RL 训练在乎确定性

强化学习的 reward 信号噪声大,如果 attention backward 本身有数值抖动,很难区分"是策略在变"还是"kernel 不确定"。确定性模式把这一层噪声抹掉,是 FA-4 对 RL 场景的针对性支持。

6. LPT 调度:架构无关的负载均衡

前面所有技术都依赖 Blackwell 专属硬件。但 FA-4 有一个完全架构无关的优化——LPT 调度,FA-3 里已有,FA-4 强化 [论文 §3.3]

问题:causal mask 和变长序列让不同 worktile 的主循环长度不同,负载不均,尾部拖慢整体 [论文 §3.3]

对你(PPU810E)的意义

LPT 调度是 FA-4 里唯一不依赖 Sm100 的技术,纯软件调度策略,理论上任何平台都能用(FA-3 已采纳)。其他 TMEM/2-CTA/DSMEM/确定性锁都是 Blackwell 专属。这条边界和第 5 课的 PPU810E caveat 一致。

7. CuTe-DSL:用 Python 写 kernel

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 [官方仓库]

8. FlashAttention 1–4 瓶颈转移主线

七节课下来,一条主线贯穿 1–4 代——每代解决上一代的瓶颈,结果瓶颈换个地方再出现

目标硬件解决的瓶颈新瓶颈
FA-1A100HBM 读写量(S/P 矩阵落地)并行度、非 matmul 开销
FA-2A100loop order、并行度、非 matmulTensor Core 空闲、softmax bubble
FA-3H100Tensor Core 利用率(WGMMA/TMA/异步重叠/FP8)exp 单元、共享内存带宽未随 MMA 翻倍
FA-4B200非对称缩放(软仿真 exp / TMEM / 2-CTA / LPT)——(截至 2026)

更完整的对比(年份、论文、精度、性能数字)见 1-4 代演进速查,术语查 术语表

主线完结

到这里,FlashAttention 1–4 代的原理与演进主线已完整覆盖。MISSION 的核心目标——"逐代比较优化重点、目标硬件和瓶颈转移"——达成。后续若感兴趣,可深入 backward recomputation 的 CUDA 实现细节,或在 PPU810E 上做实测选型专题(见 NOTES.md 待定项)。

9. 快速测验

1. FA-4 后向的主要瓶颈是什么(区别于前向的 exp)?
2. 2-CTA MMA 模式如何减少共享内存流量?
3. 确定性模式主要靠什么实现可复现的 dQ?

10. 延伸阅读

本课推荐阅读 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

理论 vs 实测

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 实测专题仍可作为后续方向。