跳到主要内容

L2.5 极致算子实战案例 (SOTA Kernel)

三维坐标 layer: L2(数据与算子)level: Seniorpillar: 编程与编译 + 训推框架

本文是 L2「数据与算子」支柱的实战收口篇。前面我们讲了算子如何被编译到硬件,本文则切到当代 LLM 性能的真正命门——注意力算子。目标不是泛泛介绍 FlashAttention,而是讲透它为什么快、快在哪个物理量上,并让你亲手用 Triton 写出一个能与 PyTorch SDPA 对齐数值的简化版 kernel。

学习目标

  • 前置知识:读过 L2 前序(尤其算子如何被编译落到硬件、GPU 存储金字塔 SRAM/HBM 的速度与容量差异);理解 L1.1 的 Tensor Core 与「访存比计算慢一个数量级」这堵物理墙;会写基础 PyTorch,知道 softmax、矩阵乘的含义即可。无需 Triton 经验。
  • 学完产出:① 能一句话说清 FlashAttention「快在哪个物理量上」——不是省 FLOPs,而是不把 N×N 分数矩阵物化到 HBM,把访存从 O(N²) 降到 O(N);② 能推导 online softmax 的 running max / running sum 递推,并解释「为什么不看整行也能正确归一化」;③ 能讲清重计算(recomputation)是如何用算力换显存、让长序列训练成为可能的;④ 能对照 FA-1/2/3 三代的核心差异,说清 FA-3 在 Hopper 上靠 TMA + wgmma 构建异步流水线的「白嫖同步」本质;⑤ 亲手用 Triton 写一个简化 FlashAttention 前向 kernel,跑通与 PyTorch SDPA 的 torch.allclose 数值对齐与速度对照,并理解教学版与工业级 kernel 的优化差距。
  • 阅读姿势:盯住一条主线——「memory-bound 算子的所有优化,本质都在让访存连续、让计算密集,尽量别让数据落回慢速 HBM」。无论是 FlashAttention 的 tiling,还是 MoE 的 grouped GEMM 分桶,解决的都是同一个问题:把零散、反复的访存,重组成一次性塞进片上 SRAM 的密集计算。

背景与现状

标准注意力的计算是 softmax(Q·Kᵀ / √d)·V。问题不在浮点运算量(FLOPs),而在那张中间的 N×N attention 分数矩阵:序列长度 N=8192 时,单头单 batch 的分数矩阵就是 8192×8192×2字节 ≈ 128MB,必须先写回 HBM(高带宽显存) 再读回来做 softmax。注意力因此是典型的 memory-bound(访存受限) 算子——GPU 的 Tensor Core 大量时间在等数据,而不是在算。

FlashAttention 的洞察一句话概括:绝不把 N×N 矩阵物化到 HBM。它用 tiling(分块) 把 Q/K/V 切成能塞进片上 SRAM 的小块,在 SRAM 里完成「分数 → softmax → 加权求和」的全过程,靠 online softmax(在线 softmax) 维护跨块的 running max / running sum,使 HBM 读写从 O(N²) 降到 O(N)。这就是它作为 「IO-aware 精确注意力」 的本质——结果和标准注意力逐位相等,省的纯粹是访存。

从产业演进看,三句话:

  • FA-1(2022):提出 tiling + online softmax + 重计算,首次让长序列训练在不掉精度的前提下提速 2–4×、省显存 10–20×。
  • FA-2(2023):重排循环顺序、减少非 matmul 的逐元素操作、把并行度铺到序列维与头维,把 A100 上的 MFU 利用率拉到 50–73%。
  • FA-3(2024):面向 Hopper(H100),用 wgmma(warpgroup 异步矩阵乘)+ TMA(Tensor Memory Accelerator 异步搬运)构建异步流水线,并引入 FP8 路径,把 H100 上的吞吐再推高一截。

业界信号Dao-AILab/flash-attention 已是几乎所有主流训练/推理框架(vLLM、Megatron-LM、TransformerEngine、PyTorch SDPA 的 flash 后端)的默认注意力内核。「会不会写、会不会调 FlashAttention」是当代算子工程师的硬通货。

原理与架构

理解 FlashAttention,必须先理解 GPU 的存储金字塔:SRAM 极快但极小(A100 每 SM 约 192KB),HBM 大但慢一个数量级。标准注意力的罪过,是把一张本可以「算完即弃」的 N×N 矩阵,硬生生在这两层之间来回搬。

2.1 标准注意力 vs FlashAttention 的数据流对比

读这张图的关键:左侧标准路径把 N×NSP 反复写读 HBM——这是 O(N²) 的访存;右侧 FlashAttention 把 Q/K/V 切块搬进 SRAM,在片上把一个行块对应的全部列块算完,只把最终的 O(N×d)和归一化用的 logsumexp(N)写回 HBM,访存降到 O(N)。两者数学结果完全一致。

2.2 online softmax:为什么能不物化整行就归一化

标准 softmax 需要看到一整行的所有分数才能算分母(求 max 防溢出、再求 exp 之和)。FlashAttention 一次只看一个列块,靠增量更新做到等价:

  • 维护 running max m(当前已见过分数的最大值)和 running sum (缩放后的 exp 之和),以及当前累加输出 O
  • 来了新列块,算出本块局部 max m_new = max(m, m_block)
  • rescale 旧状态:旧的 O 乘以 exp(m - m_new) 做对齐(因为 max 变了,之前的 exp 基准要校正)。
  • 累加本块贡献,更新 O
  • 全部列块处理完,O / ℓ 即为该行块的最终注意力输出。

这套 running max / running sum 递推,就是 FlashAttention 不物化 N×N 的数学地基。

2.3 重计算(recomputation):用算力换显存

反向传播原本需要前向存下的 N×N 概率矩阵 P——又是 O(N²) 显存。FlashAttention 选择不存 P,反向时用前向存下的 O(N) 的 logsumexp 重新算 P。这是经典的算力换显存:多花一点 FLOPs 重算,省掉 O(N²) 的激活显存,使长序列训练成为可能。

2.4 FA-1 → FA-2 → FA-3 的演进对照

维度FA-1FA-2FA-3(Hopper)
核心贡献tiling + online softmax + 重计算重排循环、减少非 matmul 操作、并行度铺到序列维异步流水线 + FP8
关键技术SRAM 分块warp 间工作划分优化wgmma 异步 MMA + TMA 异步搬运
典型硬件A100 (Ampere)A100 (Ampere)H100 (Hopper)
精度路径FP16/BF16FP16/BF16FP16/BF16 + FP8
直观收益不物化 N×N、省显存MFU 拉到 50–73%让计算与搬运重叠、吞吐再升

值得注意的是:FA-3 的核心不是「换了个公式」,而是软硬协同的流水线——TMA 异步把下一块 K/V 搬进 SRAM 的同时,wgmma 在算当前块。计算单元几乎不再空等访存,这正是 memory-bound 算子追求的终极形态。

2.5 旁支战场:量化内核与 MoE Token Dispatch

  • 量化内核(W8A16 / W4A16 / FP8)W8A16 指权重 INT8、激活 FP16;W4A16(如 AWQ/GPTQ)权重压到 4-bit,推理时在 kernel 内反量化回 FP16 再做 matmul,省的是权重的 HBM 带宽与显存FP8 则是 Hopper 起原生支持的训练/推理精度,权重激活都用 8-bit 浮点,配合 per-tensor/per-channel scale 维持数值稳定。混合精度内核的工程难点在于反量化与 matmul 的融合——不能反量化后落 HBM 再读回,否则白省。
  • MoE Top-K Token Dispatch:MoE 层每个 token 经门控选 Top-K 个专家。内核挑战是高并发的 token 路由——成千上万 token 要被分发(dispatch)到不同专家的 GEMM 上,再按原顺序聚合(combine)回来。工程上靠 grouped GEMM / 排序分桶把同一专家的 token 聚到一起做批量矩阵乘,避免稀疏访存把带宽打散。这与 FlashAttention 共享同一哲学:让访存连续、让计算密集

动手实践:用 Triton 手写简化 FlashAttention

实验目标:用 Triton 写一个简化版 FlashAttention 前向 kernel(tiling + online softmax),与 PyTorch 的 scaled_dot_product_attention(SDPA)对数值(torch.allclose)、对速度。产出物:一段 allclose=True 的数值校验输出 + 一份 Triton vs SDPA 的耗时对照。无 GPU 时,退化为「用 SDPA 的 math / flash 后端对照原理」。

3.1 环境准备

# 推荐 Python 3.11;用 uv 或 venv 隔离
python3 -m venv .venv && source .venv/bin/activate

# —— NVIDIA GPU 路径(Triton 内核需要 CUDA)——
pip install torch # 自动选 CUDA wheel
pip install triton # Linux + NVIDIA GPU 下随 torch 附带或单独装

# —— Mac / CPU 替代路径(无法跑 Triton kernel,仅跑 SDPA 对照)——
# pip install torch --index-url https://download.pytorch.org/whl/cpu

Triton kernel 只能在 NVIDIA GPU 上执行。无 GPU 的同学跳到 3.4 用 SDPA 后端对照原理,或在 Google Colab(免费 T4)上跑完整版。

3.2 NVIDIA GPU 路径:Triton 简化 FlashAttention 前向

import torch
import triton
import triton.language as tl


@triton.jit
def _flash_attn_fwd(
Q, K, V, O,
stride_qb, stride_qh, stride_qm, stride_qd,
stride_kb, stride_kh, stride_kn, stride_kd,
stride_vb, stride_vh, stride_vn, stride_vd,
stride_ob, stride_oh, stride_om, stride_od,
N, scale,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, D: tl.constexpr,
):
# 每个 program 负责一个 (batch, head) 的一个 Q 行块(BLOCK_M 行)
start_m = tl.program_id(0)
off_bh = tl.program_id(1)
off_b = off_bh // tl.num_programs(2) if False else off_bh # 简化:bh 合并维
# 基址偏移(按 batch*head 展平)
q_base = Q + off_bh * stride_qh
k_base = K + off_bh * stride_kh
v_base = V + off_bh * stride_vh
o_base = O + off_bh * stride_oh

offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_d = tl.arange(0, D)

# 载入 Q 行块到 SRAM(寄存器/共享内存由 Triton 调度)
q_ptrs = q_base + offs_m[:, None] * stride_qm + offs_d[None, :] * stride_qd
q = tl.load(q_ptrs, mask=offs_m[:, None] < N, other=0.0)

# online softmax 状态:running max m、running sum ℓ、累加输出 acc
m_i = tl.full((BLOCK_M,), float("-inf"), dtype=tl.float32)
l_i = tl.zeros((BLOCK_M,), dtype=tl.float32)
acc = tl.zeros((BLOCK_M, D), dtype=tl.float32)

# 内循环:扫过所有 K/V 列块
for start_n in range(0, N, BLOCK_N):
offs_n = start_n + tl.arange(0, BLOCK_N)
k_ptrs = k_base + offs_n[:, None] * stride_kn + offs_d[None, :] * stride_kd
v_ptrs = v_base + offs_n[:, None] * stride_vn + offs_d[None, :] * stride_vd
k = tl.load(k_ptrs, mask=offs_n[:, None] < N, other=0.0)
v = tl.load(v_ptrs, mask=offs_n[:, None] < N, other=0.0)

# 本块分数 S = scale * Q·Kᵀ
s = tl.dot(q, tl.trans(k)) * scale
s = tl.where(offs_n[None, :] < N, s, float("-inf"))

# —— online softmax 增量更新 ——
m_block = tl.max(s, axis=1)
m_new = tl.maximum(m_i, m_block)
p = tl.exp(s - m_new[:, None]) # 本块概率(未归一)
alpha = tl.exp(m_i - m_new) # 旧状态 rescale 因子
l_i = l_i * alpha + tl.sum(p, axis=1)
acc = acc * alpha[:, None] + tl.dot(p.to(v.dtype), v)
m_i = m_new

# 收尾归一化:O = acc / ℓ
acc = acc / l_i[:, None]
o_ptrs = o_base + offs_m[:, None] * stride_om + offs_d[None, :] * stride_od
tl.store(o_ptrs, acc.to(O.dtype.element_ty), mask=offs_m[:, None] < N)


def flash_attn_triton(q, k, v):
# q,k,v: [B, H, N, D],合并 B*H 为一维网格
B, H, N, D = q.shape
scale = 1.0 / (D ** 0.5)
o = torch.empty_like(q)
BLOCK_M, BLOCK_N = 64, 64
q2 = q.reshape(B * H, N, D)
k2 = k.reshape(B * H, N, D)
v2 = v.reshape(B * H, N, D)
o2 = o.reshape(B * H, N, D)
grid = (triton.cdiv(N, BLOCK_M), B * H, 1)
_flash_attn_fwd[grid](
q2, k2, v2, o2,
*q2.stride()[:1], q2.stride(0), q2.stride(1), q2.stride(2),
*k2.stride()[:1], k2.stride(0), k2.stride(1), k2.stride(2),
*v2.stride()[:1], v2.stride(0), v2.stride(1), v2.stride(2),
*o2.stride()[:1], o2.stride(0), o2.stride(1), o2.stride(2),
N, scale,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, D=D,
)
return o


if __name__ == "__main__":
assert torch.cuda.is_available(), "Triton kernel 需要 NVIDIA GPU"
torch.manual_seed(0)
B, H, N, D = 2, 4, 1024, 64
q = torch.randn(B, H, N, D, device="cuda", dtype=torch.float16)
k = torch.randn(B, H, N, D, device="cuda", dtype=torch.float16)
v = torch.randn(B, H, N, D, device="cuda", dtype=torch.float16)

o_triton = flash_attn_triton(q, k, v)
o_sdpa = torch.nn.functional.scaled_dot_product_attention(q, k, v)

# —— 对数值 ——
print("allclose:", torch.allclose(o_triton, o_sdpa, atol=1e-2, rtol=1e-2))
print("max abs diff:", (o_triton - o_sdpa).abs().max().item())

3.3 对速度

import time

def bench(fn, iters=50):
for _ in range(10): # 预热:触发 JIT / kernel 编译
fn()
torch.cuda.synchronize() # 关键:CUDA 异步,不同步等于没测
t0 = time.perf_counter()
for _ in range(iters):
fn()
torch.cuda.synchronize()
return (time.perf_counter() - t0) / iters * 1e3 # ms

t_triton = bench(lambda: flash_attn_triton(q, k, v))
t_sdpa = bench(lambda: torch.nn.functional.scaled_dot_product_attention(q, k, v))
print(f"Triton flash : {t_triton:.3f} ms")
print(f"PyTorch SDPA : {t_sdpa:.3f} ms")

你会看到 PyTorch SDPA(其底层就是优化到极致的 FlashAttention)通常更快——这是好事:它证明你的简化 kernel 数值对了,而工业级 kernel 在 tiling 大小、流水线、向量化上做了大量你这个教学版没做的优化。

3.4 Mac / CPU 替代方案:用 SDPA 后端对照原理

无 GPU 时跑不了 Triton,但能用 PyTorch SDPA 的后端切换亲手验证「FlashAttention 与朴素注意力数值一致」:

import torch
from torch.nn.attention import SDPBackend, sdpa_kernel

torch.manual_seed(0)
q = torch.randn(1, 4, 512, 64)
k = torch.randn(1, 4, 512, 64)
v = torch.randn(1, 4, 512, 64)

# math 后端:朴素实现,会物化 N×N(对照基准)
with sdpa_kernel(SDPBackend.MATH):
o_math = torch.nn.functional.scaled_dot_product_attention(q, k, v)

# flash 后端在 CPU 上通常不可用,这里用 efficient 后端作对照
with sdpa_kernel(SDPBackend.EFFICIENT_ATTENTION):
try:
o_eff = torch.nn.functional.scaled_dot_product_attention(q, k, v)
print("math vs efficient allclose:", torch.allclose(o_math, o_eff, atol=1e-4))
except RuntimeError as e:
print("该后端在当前设备不可用:", e)

目的一致:理解「同一个注意力,朴素物化 N×N 与 IO-aware 不物化,结果逐位相等、差的只是访存」。这正是 FlashAttention 的精髓。

踩坑预警 (Gotchas)

  • torch.cuda.synchronize() 就计时 = 测了个寂寞:Triton/CUDA kernel 异步下发,不同步会把耗时算成接近 0。
  • online softmax 忘了 rescale 旧状态:来了新块、max 变大后,必须把旧的 l_iaccexp(m_i - m_new),否则不同块的 exp 基准不一致,结果错得很隐蔽(小 N 看不出,大 N 数值爆炸)。
  • FP16 精度对齐用 atol=1e-2:FP16 累加误差大,用 FP32 的 1e-5 阈值会误判 allclose=False。累加器 acc/l_i 务必用 float32
  • BLOCK_M/BLOCK_N 太大导致 SRAM 溢出:Triton 会报 shared memory 超限或静默变慢。教学版用 64 稳妥,调优时按 GPU 的 SRAM 容量逐步加大。
  • stride 传错是最常见 bugtl.load 的指针算术全靠 stride,张量非连续(如 transpose 过)时必须传真实 stride,建议先 .contiguous()

深入思考

下面三题每题先给题干,再用 <details> 折叠一份图文并茂的参考答案。建议先合上答案自己想 3 分钟,再展开对照。

思考题 1:访存 vs 算力的边界

FlashAttention 把注意力从 memory-bound 改善了,但它并没有减少 FLOPs(反向甚至因重计算增加了)。结合 2.1 节的 O(N²)→O(N) 访存分析,请论证:在什么样的 N(序列长度)和 d(头维度)取值下,FlashAttention 的收益会变小甚至消失?提示——当算子从 memory-bound 滑向 compute-bound 时,省访存的边际收益如何变化?

展开参考答案(含 Roofline 收益拐点图 + 算一遍)

结论:FlashAttention 省的是访存,不是 FLOPs。它的收益完全取决于「算子有多 memory-bound」。当 d 很大、N 很小时,计算密度(算术强度)高、算子本来就接近 compute-bound,省访存的边际收益趋近于零;而 N 越大、d 越小,N×N 矩阵越主导访存、算子越 memory-bound,FlashAttention 的收益越大。

算一遍(数量级估算):注意力的 FLOPs 约 4·N²·d(QKᵀ 与 PV 各 2·N²·d),而标准实现物化 N×N 矩阵带来的 HBM 访存约 2·N² 字节(FP16)。两者的比值——算术强度 ≈ FLOPs / 访存字节 ≈ d 量级

  1. d=64N=8192(典型长序列):N×N 矩阵 = 8192² × 2B ≈ 128MB,要在 HBM 来回搬好几趟,算子深度 memory-bound——FlashAttention 把这 128MB 整个从 HBM 抹掉,提速 2–4×、省显存 10–20×。
  2. d 拉到 256、N 只有 512:N×N 矩阵 = 512² × 2B ≈ 0.5MB,而每个元素要做 d=256 次乘加,算术强度高、算子已接近 compute-bound。此时 Tensor Core 才是瓶颈,省那 0.5MB 访存对总耗时几乎无感——FlashAttention 的相对收益趋近于零
  3. 边际收益的本质:Roofline 上 memory-bound 段(左侧斜坡)每省一份访存就直接换来提速;一旦算子越过拐点进入 compute-bound 段(右侧平台),瓶颈变成算力,省访存只是「让等待的人更不忙」,不再缩短关键路径。

一句话:FlashAttention 的收益 = 算子的「memory-bound 程度」。N 越长、d 越小,收益越大;N 短、d 大到把算子推过 Roofline 拐点时,收益消失。

思考题 2:FA-3 的流水线本质

FA-3 用 TMA 异步搬运 + wgmma 异步矩阵乘实现「搬下一块的同时算这一块」。结合 2.4 节,若让你在没有 TMA 的 Ampere(A100)上模拟这种重叠,你会用哪些手段(double buffering / cp.async / 多 stage 流水线)?为什么 Hopper 的硬件原语能让这件事「白嫖」掉显式同步开销?

展开参考答案(含同步 vs 异步流水线对比图)

结论:Ampere 上要用 cp.async + double/multi-buffering 软件流水线来「手动」让访存与计算重叠——但这要写大量 stage 调度、commit_group/wait_group 同步与显式 buffer 轮换;Hopper 的 TMA + wgmma 把「异步搬运 + 异步矩阵乘」做成硬件原语,由专门的搬运引擎与 warpgroup 异步执行,程序员几乎不用写显式同步,所以叫「白嫖」掉了同步开销。

对比表:

维度Ampere(软件模拟)Hopper(硬件原语)
异步搬运cp.async:绕过寄存器把 global→shared,但仍需 warp 发射指令、占调度槽TMA:专用引擎一条指令搬整块,自动算地址,几乎不占 warp
重叠手段double / multi-buffering,手动轮换 N 个 SRAM bufferwgmma + TMA 天然异步,搬与算分属不同硬件单元
同步成本显式 cp.async.commit_group / wait_group,程序员管 stage 依赖mbarrier 硬件 barrier,软件几乎不写显式 wait
寄存器压力高:软件流水要在寄存器里维持多 stage 状态低:异步引擎自带状态,warp 更省

为什么 Hopper 能「白嫖」:Ampere 的重叠是软件假象——warp 仍要花指令去发起 cp.async、去 wait_group,这些都占用宝贵的调度槽和寄存器,只是把「等访存」藏进了流水线缝隙。Hopper 把搬运(TMA)和矩阵乘(wgmma)拆成物理上独立、各自异步的硬件单元,再用 mbarrier 在硬件层面协调依赖——程序员只需「发起 → 等屏障」,不必手写多 stage 调度。这与 L1.1 讲的昇腾「显式流水编排」殊途同归:让搬运与计算分属不同单元、在时间轴上完全重叠,区别只在于 Hopper 把这层编排更多地下沉到了硬件。

思考题 3:MoE 与注意力的访存哲学统一

FlashAttention 靠 tiling 让访存连续,MoE 的 Top-K dispatch 靠 grouped GEMM / 排序分桶让访存连续。结合 2.5 节,请用「让访存连续、让计算密集」这一条主线,解释为什么稀疏 MoE 的工程难点本质上不是稀疏计算,而是稀疏访存的重新聚合

展开参考答案(含 MoE token 路由分桶聚合图)

结论:MoE 的「稀疏」体现在「每个 token 只激活 Top-K 个专家」,但每个被选中的专家内部做的仍是稠密 GEMM——计算本身一点都不稀疏。真正的难点是:门控把 token 打散到了不同专家,若直接按 token 顺序逐个去对应专家的权重做乘法,访存就变成零散、跳跃的随机访问,把 HBM 带宽彻底打碎。所以工程要做的是把同一专家的 token 重新聚合到一起(排序分桶 → grouped GEMM),让访存重新连续——这和 FlashAttention 用 tiling 让访存连续是同一条哲学。

对比一遍:注意力 vs MoE 的「连续化」手段

维度FlashAttentionMoE Top-K Dispatch
稀疏/零散的来源N×N 分数矩阵太大,逐元素搬 HBM 是 O(N²) 零散访存门控把 token 打散到不同专家,逐 token 取权重是随机访存
连续化手段tiling:Q/K/V 切块塞进 SRAM,块内连续算完排序分桶 + grouped GEMM:同专家 token 聚到一起批量乘
计算本身稠密 matmul,不变每个专家内部仍是稠密 GEMM,不变
省的是什么HBM 访存(O(N²)→O(N))HBM 带宽(避免随机访存把带宽打散)

为什么难点不是稀疏计算:MoE 的 FLOPs 其实不大——每个 token 只过 K 个专家,激活的计算量远小于稠密 FFN。如果瓶颈在计算,那 MoE 应该「天生就快」。但现实是 MoE kernel 的工程量巨大,原因全在访存重组:门控输出是一份「token→专家」的乱序映射,要先 dispatch(把同专家 token 物理聚到连续内存)、做 grouped GEMM、再 combine(按原顺序散回原位)。这两次 scatter/gather 才是真正吃带宽、吃工程量的地方。再叠加负载均衡(防止热门专家拥堵、冷门专家闲置)与专家并行下的 all-to-all 通信,难点全部围绕「如何把稀疏、跳跃的访存,重新聚合成连续、密集的访存与计算」——和 FlashAttention 用 tiling 把 O(N²) 零散访存压成 O(N) 连续访存,是同一条主线的两个战场

延伸阅读

1. 核心 Paper

  • FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(Dao et al., 2022)— 必读,tiling + online softmax + 重计算的奠基之作。
  • FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning(Dao, 2023)— 理解 FA-2 如何把并行度铺满、把非 matmul 操作压到最少。
  • FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision(2024)— 理解 Hopper 的 wgmma/TMA 异步流水线与 FP8 路径。
  • Online normalizer calculation for softmax(Milakov & Gimelshein, 2018)— online softmax 的数学源头。

2. 相关高 Star 仓库与源码必读路径

  • Dao-AILab/flash-attention — 工业级实现,看 csrc/flash_attn/ 与 Python 封装,对照本文教学版理解优化差距。
  • triton-lang/triton — 看 python/tutorials/06-fused-attention.py,官方的 Triton FlashAttention 教程(含反向),是本文实验的进阶版。
  • vllm-project/vllm — 看 vllm/attention/ 后端选择,理解 FlashAttention 如何被推理引擎集成;MoE 看 vllm/model_executor/layers/fused_moe/

3. 优质博客 / 视频

  • Tri Dao 的 FlashAttention 系列官方博客与讲座,原作者视角讲清设计取舍。
  • Triton 官方文档「Fused Attention」教程,逐行讲解 kernel 写法与 autotune。
  • NVIDIA「Hopper Architecture」白皮书相关章节,理解 TMA / wgmma / FP8 的硬件原语。

下一篇L2.6 算子 Shape 泛化与性能分析:把视角从「写一个 SOTA kernel」放大到「让 kernel 在各种 shape 下都不掉性能」,讲清形变泛化、autotune 与系统级 profiling。