Attention Kernel:FlashAttention 与它的后继者
大模型时代最重要的一次 kernel 重写:不改数学,只改数据从哪来——attention 的显存占用从
降到 ,速度还更快。
问题:标准 attention 卡在哪
标准实现分三步:
运算强度算下来:
FlashAttention 的洞察:IO-aware
FlashAttention(Dao et al., 2022)没有改数学,改的是遍历顺序:
- 分块(tiling):把
、 、 切成小块, 块逐块流入 SRAM(shared memory),在 SRAM 内完成"局部 → 局部 softmax → 累积输出"。 - 在线 softmax(online softmax):softmax 的分母需要整行 max/sum,而分块后每块只见局部数据。解法是维护运行统计量(当前行最大值
与和 ),每来一个新块就重缩放已累积的输出: 。数学上精确等价,物理上从不需要完整的一行。
效果:
FlashAttention-2 与 -3:往硬件上限逼近
- FA-2(2023):重排了循环——外层遍历
块、内层遍历 块,让同一输出块的多次累积留在寄存器;并按 warp 重分任务减少非矩阵乘指令。吞吐比 FA-1 再翻近一倍,"接近 GEMM 的 attention"。 - FA-3(2024,Hopper):用 TMA 异步搬运 + warp 专业化(生产者-消费者 warp 组)+ ping-pong 调度(softmax 与 GEMM 在不同 warp 上交错),把 attention 推到 Hopper Tensor Core 峰值的 75% 量级——本质是把算子优化篇的异步流水技术用到极致。
一句话概括这条演进线:FA-1 赢在算法(IO 复杂度),FA-2 赢在并行度,FA-3 赢在硬件流水——每一代都先把上一代的理论下限变成新的起点。
生态位:xFormers、cuDNN 与稀疏 attention
- xFormers:memory-efficient attention 的早期集大成者(FA 同期的另一实现),以
memory_efficient_attentionAPI 提供多种 kernel 后端,适合作为 PyTorch 生态的兼容层;如今多数场景被 FA 系列取代。 - cuDNN fused attention:NVIDIA 官方的融合 attention,Hopper 上与 FA-3 性能相当,优点是框架集成(
torch.nn.functional.scaled_dot_product_attention的 cudnn 后端)与更广的 shape 覆盖。 - 稀疏 attention:把
里不重要的块跳过(滑窗、块稀疏、hyper-attention 类)。它是唯一能突破 计算量的路线,但kernel 实现受制于不规则的访存模式(gather 打破 coalescing、负载不均),且"哪些块重要"的近似有精度风险——工程上先确认真的被长序列计算卡住,再考虑稀疏(见[芯片架构]篇"先算运算强度再谈优化")。
为什么 FA 的分块对训练同样成立(反向传播怎么办)
反向传播需要前向的中间量(softmax 输出),而 FA 没存它们。解法:重计算——反向时用保存的运行统计量(
小结
- 标准 attention 的瓶颈不是计算而是
中间矩阵的 HBM 往返;FlashAttention 用分块 + 在线 softmax 把它压进 SRAM。 - FA-1 算法、FA-2 并行、FA-3 硬件流水;每代的共同量尺是"离 GEMM 上限还有多远"。
- 生态上 FA 系列是主力,cuDNN fused attention 是官方集成选项,xFormers 是兼容层。
- 稀疏 attention 是突破平方复杂度的唯一路线,但工程与精度代价要求先测量再采用。
思考题
、 、32 头、BF16:标准实现的 矩阵多大?一张 80 GB 的卡装得下吗? - 在线 softmax 为什么数学上等价?写出两块的累积更新式。
- 你的模型 prefill 4K 序列时 attention 只占 8% 时间,值得换 FA-3 吗?
参考答案
(32 头 × × 2 字节)——单卡装不下,这正是长序列必须 FA/分块的定量理由。 - 设两块的局部 max 为
、指数和为 、部分输出为 ;合并时 , , ——最终归一化时 恰好抵消,与整体 softmax 精确一致。 - 不值得。FA-3 的收益集中在长序列(attention 占比高的场景);4K 时 attention 占比 8%,即使快一倍也只省 4% 总时间,优先去优化占比更大的项(GEMM、通信、采样)。
参考资料
- Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(arXiv 2205.14135)
- Dao, FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning(arXiv 2307.08691)
- Shah et al., FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision(arXiv 2407.08608)
- NVIDIA, cuDNN scaled dot product attention 文档