Skip to main content

第25讲:为什么算 Attention 可以不存中间矩阵?——FlashAttention 原理剖析与 Triton 工业级实战

主讲人:👓 Ringi(大厂 AI Infrastructure 资深性能架构师)
所属专栏:《AI_Infra大话西游之水滴石穿》 ➔ Module 02: CUDA 编程与高性能算子优化
篇章范式:⚡ 体系结构感知算子工程范式(Architecture-Aware Operator & Kernel Engineering Paradigm)
源码与实验环境:NVIDIA A100-SXM4-80GB / H100-SXM5-80GB | CUDA 12.4 | Python 3.10 | PyTorch 2.3+ | Triton 2.3+
知识底账索引:
  • 核心理论证据:FlashAttention 深度剖析(AI-fundamentals)
  • 原理与公式证明:FlashAttention V1 详解(AIInfraGuide)
  • 经典算子源码:LeetCUDA 算子实战(LeetCUDA)

Ringi 导师解构:FlashAttention 突破存储墙全景工坊

0. Ringi 开场:生产真实现场与痛点冲突

0.1 真实工程矛盾:128K 上下文下的“中间矩阵显存血崩”

各位做大模型系统架构的同袍,在日常的线上支持和性能调优中,你一定被算法同学追问过类似的问题:
“Ringi,我们买的明明是单卡 80GB 显存的 A100/H100,为什么预训练或者推理 Batch Size 设成 1,仅仅把上下文长度(Sequence Length)拉到 128K,显卡瞬间就报 CUDA out of memory?我们的模型权重明明才 14GB 啊,剩下 66GB 的显存到底被什么幽灵给吞了?!”
答案极其简单,但也极其残酷:吞噬显存的根本不是模型权重,而是注意力机制中那个 N×NN \times N 的中间注意力矩阵(Attention Map)! 让我们掏出工程算盘手算一笔账: 标准的多头注意力(Multi-Head Attention)包含三个核心矩阵运算:
  1. S=QKTS = Q K^T
  2. P=softmax(S)P = \text{softmax}(S)
  3. O=PVO = P V
当输入序列长度 N=131,072N = 131,072(128K),Head 维度 d=128d = 128,数据类型为 FP16(2 字节/元素)时: 单个 Attention Head 的注意力得分矩阵 SS 和归一化概率矩阵 PP 的尺寸是: N×N=131,072×131,072≈1.718×1010 个元素N \times N = 131,072 \times 131,072 \approx 1.718 \times 10^{10} \text{ 个元素} 仅仅存下矩阵 SS,就需要: 1.718×1010×2 Bytes≈34.36 GB1.718 \times 10^{10} \times 2 \text{ Bytes} \approx 34.36 \text{ GB} 而计算完 Softmax 之后得到的概率矩阵 PP,同样是 N×NN \times N,又需要 34.36 GB! 这意味着,光是算一个 Attention Head,就需要 68.7 GB 的物理显存来存放这两个临时中间矩阵! 如果你的模型有 32 个 Query Head(即便有 GQA/MQA),如果用 PyTorch 原生的三步实现,哪怕只存一个 Head 的激活值,80GB 的显存就已经宣告爆仓熔断。 更致命的是访存带宽。在 A100 上,HBM2e 的物理带宽极限是 2.0 TB/s(实际有效利用率约 1.5 TB/s)。将这 68.7 GB 的数据写出到 HBM,再读取出来传给 VV 做乘法: 访存耗时=68.7×2 GB1500 GB/s≈91.6 ms\text{访存耗时} = \frac{68.7 \times 2 \text{ GB}}{1500 \text{ GB/s}} \approx 91.6 \text{ ms} 而 A100 的 Tensor Core 算力高达 312 TFLOPS,执行这些矩阵乘法本身的纯计算耗时只要不到 5 ms。 95% 以上的时间,整张显卡的高性能 Tensor Core 都在干瞪眼,全卡都在为 HBM 漫长的数据搬运排队买单! 这就是大模型工程中最典型的“存储墙(Memory Wall)死局”。

0.2 线上真实事故复盘:某长文本对话系统引发的集群级联 OOM 熔断

2024 年初,国内某头部大模型团队在将线上长文本对话服务从 8K 灰度推向 32K 时,发生了一起严重的 P0 级线上雪崩事故。 事故现场还原: 算法团队为了支持一种特殊的相对位置编码,在代码库中绕过了 Triton / FlashAttention 内核,使用原生 PyTorch 算子手写了 Attention 过程(torch.baddbmm + torch.softmax + torch.bmm)。在 4K 和 8K 压测时,由于单卡显存能够容纳中间矩阵,测试集延迟表现尚可接受。 然而,当全量流量切入,线上请求涌入大量 32K 的超长 PDF 分析任务时:
  1. 显存阶跃爆炸: 32K32\text{K} 相比 8K8\text{K},序列长度增加 4 倍,中间矩阵 SS 和 PP 的显存占用直接暴增 42=164^2 = 16 倍!
  2. CUDA 显存分配器锁死:PyTorch 的 caching_allocator 在面对单次超过 40GB 的瞬时巨型张量申请时,触发了显存碎片的紧急整理与系统级垃圾回收,导致 GPU 工作线程陷入长达数百毫秒的软锁死;
  3. 节点级联超时崩溃:显卡被巨量 HBM 搬运堵死,导致推理网关的心跳包超时,Kubernetes 集群将正在处理任务的 Pod 判定为 Unhealthy 并强制杀进程重启;而重启后流量重新路由至邻近节点,瞬间将邻近节点也打入 OOM 循环,造成了多达 64 台 8 卡 H800 服务器的连环崩溃。
事后复盘时,值班工程师唯一的修复方案,就是在 2 个小时内将底层的 Attention 算子全量回滚并熔接到 FlashAttention-2。回滚完成后,显存占用瞬间暴跌 90%,单步延迟缩短 4.2 倍,集群负载曲线瞬间平稳如镜。

0.3 AI Infra 注意力算子演进全景速查表

在深入数学推导与代码之前,我们先拉出一张现代 AI 基础设施中关于 Attention 算子演进的高阶对照表,理清每一代突破的技术本质:

1. 注意力机制的算术强度与存储墙死局

为了在深入具体细节前建立完整的物理心智模型,下方给出了标准 Attention 显存存储墙、IO 感知 SRAM 分块、Online Softmax 动态缩放状态机与 FlashAttention-1/2/3 代际演进的工业级全景架构拓扑: FlashAttention IO 感知计算、Online Softmax 融合与 Triton 编译全景图

1.1 Standard Attention 算法回顾与访存拆解

标准多头自注意力机制(Self-Attention)的数学表达式妇孺皆知: Attention(Q,K,V)=softmax(QKTd)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{Q K^T}{\sqrt{d}}\right) V 其中:
  • Q∈RN×dQ \in \mathbb{R}^{N \times d}(Query 矩阵)
  • K∈RN×dK \in \mathbb{R}^{N \times d}(Key 矩阵)
  • V∈RN×dV \in \mathbb{R}^{N \times d}(Value 矩阵)
  • NN 为序列长度(Sequence Length), dd 为头维度(Head Dimension,通常为 64 或 128)。
在传统的深度学习框架(如原生 PyTorch、TensorFlow)中,这个计算图被拆分为三个物理上独立的 CUDA Kernel 发射到 GPU 执行:

1.2 访存账本与 O(N2)O(N^2) 显存灾难

让我们仔细清点这三个独立 Kernel 在 GPU HBM 显存总线上产生的真实物理读写字节数(假设采用 FP16,每个数值 2 字节):
  1. Kernel 1( S=QKTS = Q K^T ):
    • 读 QQ: 2Nd2 N d 字节;
    • 读 KK: 2Nd2 N d 字节;
    • 写 SS: 2N22 N^2 字节;
    • 浮点运算量(FLOPs): 2N2d2 N^2 d(乘加各一次)。
  2. Kernel 2( P=softmax(S)P = \text{softmax}(S) ):
    • 读 SS: 2N22 N^2 字节;
    • 写 PP: 2N22 N^2 字节;
    • 浮点运算量(FLOPs):约 3N23 N^2(减最大值、取指数、累加求和、除法归一化)。
  3. Kernel 3( O=PVO = P V ):
    • 读 PP: 2N22 N^2 字节;
    • 读 VV: 2Nd2 N d 字节;
    • 写 OO: 2Nd2 N d 字节;
    • 浮点运算量(FLOPs): 2N2d2 N^2 d。
全流程总物理访存量(HBM IO Traffic): Total HBM Bytes=(2Nd+2Nd+2N2)+(2N2+2N2)+(2N2+2Nd+2Nd)=6N2+8Nd Bytes\text{Total HBM Bytes} = (2Nd + 2Nd + 2N^2) + (2N^2 + 2N^2) + (2N^2 + 2Nd + 2Nd) = 6 N^2 + 8 N d \text{ Bytes} 全流程总计算量(FLOPs): Total FLOPs=2N2d+3N2+2N2d≈4N2d FLOPs\text{Total FLOPs} = 2 N^2 d + 3 N^2 + 2 N^2 d \approx 4 N^2 d \text{ FLOPs}

1.3 No Naked Formula 2.0:算术强度与 Roofline 极限穿透

① 为什么算?

我们要精确量化标准 Attention 在现代 GPU 硬件微架构上到底是不是 Memory-bound(受限于访存带宽),以及硬件 Tensor Core 究竟闲置了多少。

② Mental Model(炒菜装盘比喻)

想象你是一个顶级厨师(Tensor Core,翻炒速度极快,每秒翻炒 300 次)。
  • 标准 Attention 模式:你炒完了半成品( SS 矩阵),非要盛进盘子里,让服务员端回地下的冷库大仓库(HBM);过了一秒钟,你让服务员再从冷库把这盘半成品端回厨房操作台,淋上酱汁(Softmax 得到 PP 矩阵),然后再盛进盘子端回冷库;又过了一秒,你再让服务员把盘子从冷库端出来,和配料( VV 矩阵)一起下锅翻炒出成品( OO 矩阵)。
  • 结果:厨师 90% 的时间在等服务员跑腿端盘子,冷库大门(显存接口)被挤得水泄不通,而炉灶(Tensor Core)全程熄火等待!

③ Tiny Calculator(极简数字小算盘)

设序列长度 N=4096N = 4096,Head 维度 d=128d = 128:
  • 浮点计算量:
FLOPs=4×(4096)2×128=4×16,777,216×128≈8.59×109 FLOPs=8.59 GFLOPs\text{FLOPs} = 4 \times (4096)^2 \times 128 = 4 \times 16,777,216 \times 128 \approx 8.59 \times 10^9 \text{ FLOPs} = 8.59 \text{ GFLOPs}
  • HBM 物理访存量:
Bytes=6×(4096)2+8×4096×128=6×16.78M+4.19M≈104.86 MB\text{Bytes} = 6 \times (4096)^2 + 8 \times 4096 \times 128 = 6 \times 16.78\text{M} + 4.19\text{M} \approx 104.86 \text{ MB}
  • 算术强度(Arithmetic Intensity):
I=FLOPsBytes=8.59×109104.86×106≈81.9 FLOPs/ByteI = \frac{\text{FLOPs}}{\text{Bytes}} = \frac{8.59 \times 10^9}{104.86 \times 10^6} \approx 81.9 \text{ FLOPs/Byte} 如果序列进一步拉长,例如 N=16,384N = 16,384:
  • N2N^2 项彻底主导分母:
I=4N2d6N2+8Nd≈4d6=23d=23×128≈85.3 FLOPs/ByteI = \frac{4 N^2 d}{6 N^2 + 8 N d} \approx \frac{4 d}{6} = \frac{2}{3} d = \frac{2}{3} \times 128 \approx 85.3 \text{ FLOPs/Byte}

④ Formal Model(数学模型证明)

标准 Attention 的渐近算术强度上限为: lim⁡N→∞Istandard=23d FLOPs/Byte\lim_{N \to \infty} I_{\text{standard}} = \frac{2}{3} d \text{ FLOPs/Byte} 无论你把序列拉到多长( N=16K,64K,128KN=16\text{K}, 64\text{K}, 128\text{K} ),标准 Attention 的算术强度永远被头维度 dd 钉死在上界!当 d=64d=64 时,强度仅为 42.6 FLOPs/Byte;当 d=128d=128 时,强度仅为 85.3 FLOPs/Byte。

⑤ Sanity Check(A100 硬件真实物理校验)

看一下 NVIDIA A100 SXM4 的硬件指标:
  • 半精度 Tensor Core 峰值算力:
Cpeak=312 TFLOPS=3.12×1014 FLOPs/sC_{\text{peak}} = 312 \text{ TFLOPS} = 3.12 \times 10^{14} \text{ FLOPs/s}
  • HBM2e 物理峰值带宽:
Bpeak=2.039 TB/s=2.039×1012 Bytes/sB_{\text{peak}} = 2.039 \text{ TB/s} = 2.039 \times 10^{12} \text{ Bytes/s}
  • 硬件拐点算术强度(Roofline Knee):
Iknee=CpeakBpeak=312×10122.039×1012≈153 FLOPs/ByteI_{\text{knee}} = \frac{C_{\text{peak}}}{B_{\text{peak}}} = \frac{312 \times 10^{12}}{2.039 \times 10^{12}} \approx 153 \text{ FLOPs/Byte} 物理结论一目了然: 硬件要求每个字节的访存必须支撑至少 153 次浮点计算,才能让 Tensor Core 完全满载! 而 Standard Attention 无论序列多长,算术强度最多只有 85 FLOPs/Byte( d=128d=128 )甚至 42 FLOPs/Byte( d=64d=64 )。 它从物理层面上就被判了死刑——永远死死卡在 Roofline 模型的 Memory-bound 倾斜上升段!你的 Tensor Core 无论如何优化代码,利用率理论上限也绝对超不过 55%!

2. Online Softmax 算法数学推导与 2D 分块(Tiling)

要想打破这个死局,思路非常明确:能不能把中间的 N×NN \times N 矩阵完全留在芯片内部极速的 SRAM(Shared Memory)里,根本不写出到外部 HBM? 但立即会遇到一个看似无法逾越的数学鸿沟:Softmax 的全归一化依赖!

2.1 传统 Safe Softmax 的两次全局遍历困境

给定一个长度为 NN 的向量 x=[x1,x2,…,xN]x = [x_1, x_2, \dots, x_N],为了防止浮点数指数运算发生上溢( exi→∞e^{x_i} \to \infty ),工业界通行的 Safe Softmax 必须包含三步操作:
  1. 求全局最大值:
m=max⁡1≤k≤Nxkm = \max_{1 \le k \le N} x_k
  1. 计算平移指数并累加全局配分函数(分母):
ℓ=∑k=1Nexk−m\ell = \sum_{k=1}^N e^{x_k - m}
  1. 计算每个元素的归一化概率:
pi=exi−mℓp_i = \frac{e^{x_i - m}}{\ell} 在传统的体系结构中,这需要两次全局数据遍历:
  • 第一遍(Pass 1):扫描全部 NN 个元素,算出最大值 mm;
  • 第二遍(Pass 2):再次扫描全部 NN 个元素,算出分母 ℓ\ell 并完成除法归一化。
如果你把数据切分成多个 Block 分别加载到 SRAM 里,当你在处理第一个 Block 时,你根本不知道后面还没读进来的 Block 里有没有比当前更大的数值!你不知道全局的 mm,更不知道全局的 ℓ\ell,你如何能在 SRAM 内部把局部结果乘以 VV 算出来?!

2.2 Online Softmax 递推公式的严密数学推导

FlashAttention 的核心数学基石,正是由 Milakov & Gimelshein 在 2018 年提出、并在 FlashAttention 中发扬光大的 Online Softmax(流式在线 Softmax) 算法。 它证明了一个极其优美的数学结论:只要引入一个尺度修正因子(Rescaling Factor),我们可以在单遍扫描的过程中,一边流式读入新分块,一边动态修正之前已经累加的历史输出! 让我们给出无跳步的严密推导: 假设一个长向量 xx 被切分为两段(或我们已经处理了前 k−1k-1 个元素,当前新读入了第 kk 个元素):
  • 设历史段(Previous)的最大值为 mprevm_{\text{prev}},配分函数和为 ℓprev=∑i∈prevexi−mprev\ell_{\text{prev}} = \sum_{i \in \text{prev}} e^{x_i - m_{\text{prev}}};
  • 当前新段(Current)的最大值为 mcurrm_{\text{curr}},配分函数和为 ℓcurr=∑j∈currexj−mcurr\ell_{\text{curr}} = \sum_{j \in \text{curr}} e^{x_j - m_{\text{curr}}}。
现在我们合并这两段,新的全局联合最大值 mnewm_{\text{new}} 显然为: mnew=max⁡(mprev,mcurr)m_{\text{new}} = \max(m_{\text{prev}}, m_{\text{curr}}) 那么,真正的联合配分函数 ℓnew=∑t∈prev∪currext−mnew\ell_{\text{new}} = \sum_{t \in \text{prev} \cup \text{curr}} e^{x_t - m_{\text{new}}} 该如何通过局部的 ℓprev\ell_{\text{prev}} 和 ℓcurr\ell_{\text{curr}} 得到? 我们展开推导: ℓnew=∑i∈prevexi−mnew+∑j∈currexj−mnew=∑i∈preve(xi−mprev)+(mprev−mnew)+∑j∈curre(xj−mcurr)+(mcurr−mnew)=emprev−mnew∑i∈prevexi−mprev+emcurr−mnew∑j∈currexj−mcurr=ℓprev⋅emprev−mnew+ℓcurr⋅emcurr−mnew\begin{aligned} \ell_{\text{new}} &= \sum_{i \in \text{prev}} e^{x_i - m_{\text{new}}} + \sum_{j \in \text{curr}} e^{x_j - m_{\text{new}}} \\ &= \sum_{i \in \text{prev}} e^{(x_i - m_{\text{prev}}) + (m_{\text{prev}} - m_{\text{new}})} + \sum_{j \in \text{curr}} e^{(x_j - m_{\text{curr}}) + (m_{\text{curr}} - m_{\text{new}})} \\ &= e^{m_{\text{prev}} - m_{\text{new}}} \sum_{i \in \text{prev}} e^{x_i - m_{\text{prev}}} + e^{m_{\text{curr}} - m_{\text{new}}} \sum_{j \in \text{curr}} e^{x_j - m_{\text{curr}}} \\ &= \ell_{\text{prev}} \cdot e^{m_{\text{prev}} - m_{\text{new}}} + \ell_{\text{curr}} \cdot e^{m_{\text{curr}} - m_{\text{new}}} \end{aligned} 看!这是一个何等精妙的递推关系式! 只要我们将历史的累加和 ℓprev\ell_{\text{prev}} 乘以修正衰减因子 emprev−mnewe^{m_{\text{prev}} - m_{\text{new}}},再加上新分块的修正和,就严丝合缝、毫无数值损失地得到了新的全局配分函数!由于 mnew≥mprevm_{\text{new}} \ge m_{\text{prev}},指数上的差值 mprev−mnew≤0m_{\text{prev}} - m_{\text{new}} \le 0,指数运算的结果恒在 (0,1](0, 1] 之间,永远不会发生浮点上溢!
Ringi 导师解构:Online Softmax 动态 Rescaling 流水线

2.3 2D Tiling 矩阵乘的输出动态 Rescaling 机制

现在,我们把这个标量推导推广到真正的 Attention 矩阵乘法: O=Softmax(QKT)VO = \text{Softmax}(Q K^T) V 设输出向量(或矩阵行)为 OO。 在没有完成全部块计算前,假设我们只算到了第 kk 个 Key-Value 块,当前维护的未完全归一化加权累加值为: Ounnorm(k)=∑j=1keSj−mkVjO_{\text{unnorm}}^{(k)} = \sum_{j=1}^k e^{S_j - m_k} V_j 真实的归一化输出应为: O(k)=Ounnorm(k)ℓkO^{(k)} = \frac{O_{\text{unnorm}}^{(k)}}{\ell_k} 当第 k+1k+1 个分块到来时,新的最大值为 mk+1=max⁡(mk,mcurr)m_{k+1} = \max(m_k, m_{\text{curr}})。 历史累加值 Ounnorm(k)O_{\text{unnorm}}^{(k)} 的指数基准是 mkm_k,为了将它的基准统一修正到 mk+1m_{k+1},我们必须将历史累加值整体乘以补偿因子: emk−mk+1e^{m_k - m_{k+1}} 于是,更新后的未归一化累加值为: Ounnorm(k+1)=Ounnorm(k)⋅emk−mk+1+eScurr−mk+1VcurrO_{\text{unnorm}}^{(k+1)} = O_{\text{unnorm}}^{(k)} \cdot e^{m_k - m_{k+1}} + e^{S_{\text{curr}} - m_{k+1}} V_{\text{curr}} 而新的分母为: ℓk+1=ℓk⋅emk−mk+1+∑eScurr−mk+1\ell_{k+1} = \ell_k \cdot e^{m_k - m_{k+1}} + \sum e^{S_{\text{curr}} - m_{k+1}} 这就是 Tiled Attention 能够在片上 SRAM 完成所有矩阵乘法的底层数学真理! 我们根本不需要在内存里保留完整的 SS 或 PP 矩阵!只要分配几个小寄存器,动态追踪维护当前的 (m,ℓ,O)(m, \ell, O),每算完一个子块就原地 rescale 一次累加器,当所有分块循环完毕后,做一次除法 O=Ounnorm/ℓO = O_{\text{unnorm}} / \ell,就能输出与标准 Attention 在数学上严格等价、但显存开销整整降了一个量级的最终结果!

3. FlashAttention-1 架构原理与计算流解构

在理解了 Online Softmax 之后,Tri Dao 等人在 2022 年发表的里程碑论文《FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness》的架构蓝图就呼之欲出了。

3.1 IO-Awareness 体系结构哲学

计算机体系结构领域有一句名言:“FLOPs are cheap, IO is expensive.”(算力很便宜,搬运极昂贵)。 在现代 GPU 上,由于制程技术的进步,算力提升的速度远远超过了外部物理引脚和内存总线带宽的提升速度。 FlashAttention 的核心哲学就是:宁可多花算力在片上做重复计算,也绝对不把中间数据写到慢速 HBM 上走一遭! 通过在 GPU 片上 SRAM(Shared Memory)与片外 HBM 之间建立严格的显式流水线分块,算子在 SRAM 内部一次性完成 QKTQ K^T、Online Softmax 和 PVP V 的全套闭环。

3.2 经典前向分块算法(Algorithm 1 深度解构)

在 FlashAttention-1 的原始设计中,分块逻辑如下: 设 SRAM 大小为 MM 字节。我们将输入分块:
  • Block 尺寸: Br=⌈M4d⌉B_r = \lceil \frac{M}{4d} \rceil, Bc=⌈M4d⌉B_c = \lceil \frac{M}{4d} \rceil
  • QQ 被切分成 Tr=⌈N/Br⌉T_r = \lceil N / B_r \rceil 个块: Q1,Q2,…,QTrQ_1, Q_2, \dots, Q_{T_r}
  • K,VK, V 被切分成 Tc=⌈N/Bc⌉T_c = \lceil N / B_c \rceil 个块: K1,K2,…,KTcK_1, K_2, \dots, K_{T_c} 与 V1,V2,…,VTcV_1, V_2, \dots, V_{T_c}
FlashAttention-1 的循环调度嵌套如下:

3.3 反向传播的重计算(Backward Recomputation)哲学:以算力换带宽

在深度学习反向传播中,按照标准微积分链式法则,为了计算关于 Q,K,VQ, K, V 的梯度 ∇Q,∇K,∇V\nabla_Q, \nabla_K, \nabla_V,必须用到前向传播时的注意力权重矩阵 P=softmax(QKT)P = \text{softmax}(Q K^T): ∇V=PT∇O\nabla V = P^T \nabla O ∇P=∇OVT\nabla P = \nabla O V^T ∇S=P∘(∇P−rowsum(∇P∘P))\nabla S = P \circ (\nabla P - \text{rowsum}(\nabla P \circ P)) ∇Q=∇SK,∇K=∇STQ\nabla Q = \nabla S K, \quad \nabla K = \nabla S^T Q 在标准实现中,为了给反向传播使用,前向传播必须把尺寸为 N×NN \times N 的 PP 矩阵完整保存在显存中!这正是导致训练阶段长上下文直接 OOM 的罪魁祸首! FlashAttention 给出了一个极具颠覆性的工程决策: 前向传播结束时,把整个 PP 矩阵彻底丢弃!显存里只保存微不足道的行统计量:标量向量 m∈RNm \in \mathbb{R}^N 和 ℓ∈RN\ell \in \mathbb{R}^N! 在反向传播执行到 Attention 算子时:
  1. 再次把对应分块的 QiQ_i 和 KjK_j 从 HBM 读进 SRAM;
  2. 利用保存的标量 mim_i 和 ℓi\ell_i,在 SRAM 内部当场重新计算一遍分块注意力概率矩阵 PijP_{ij}!
  3. 立即参与反向梯度的乘加运算!
这笔账划算吗?我们算一算:
  • 多付出的代价:多做了一次小分块矩阵乘法 QiKjTQ_i K_j^T,大约增加了 15%~20% 的浮点运算量;
  • 获得的巨大收益:
    • 激活值显存占用从 O(N2)O(N^2) 骤降到 O(N)O(N)(对于 N=64KN=64\text{K},显存从 32GB 骤降到几百 KB);
    • 消除了将 N×NN \times N 矩阵写入 HBM、再在反向传播中从 HBM 读回的庞大总线流量! 在现代 GPU 上,由于反向传播省去了海量 HBM 搬运开销,即使多算了矩阵乘法,整体反向传播的运行速度反而比原版快了 2~3 倍!

3.4 IO 复杂度严格证明(Theorem 2)

FlashAttention-1 论文中最核心的理论贡献,是证明了其 HBM 访问复杂度的下界。 定理 2(IO 复杂度): 设输入序列长度为 NN,Head 维度为 dd,GPU 片上高速缓存(Shared Memory)容量为 MM 字节。
  • 标准 Attention 的 HBM 访问总量为:
IOstandard=Θ(Nd+N2)\text{IO}_{\text{standard}} = \Theta(N d + N^2)
  • FlashAttention 的 HBM 访问总量为:
IOflash=Θ(N2d2M)\text{IO}_{\text{flash}} = \Theta\left(\frac{N^2 d^2}{M}\right) 证明简述: 在 FlashAttention 中, K,VK, V 的 Block 大小为 Bc≈M4dB_c \approx \frac{M}{4d}。
  • 外层循环遍历 K,VK, V 分块,共需要迭代 Tc=NBc=4NdMT_c = \frac{N}{B_c} = \frac{4 N d}{M} 次;
  • 在每一次外层循环中,内层循环必须遍历一遍完整的 QQ 矩阵(大小为 N×dN \times d ),因此读取 QQ 的总量为:
Read(Q)=Tc×(Nd)=4NdM×Nd=4N2d2M\text{Read}(Q) = T_c \times (N d) = \frac{4 N d}{M} \times N d = \frac{4 N^2 d^2}{M}
  • K,VK, V 在外层循环中只被加载一次,总量为 2Nd2 N d;
  • 最终输出 OO 的读写总量为 O(Nd)O(N d)。
  • 综合各项,总 HBM 访问字节数为:
IOflash=Θ(N2d2M+Nd)=Θ(N2d2M)\text{IO}_{\text{flash}} = \Theta\left(\frac{N^2 d^2}{M} + N d\right) = \Theta\left(\frac{N^2 d^2}{M}\right) 物理比率分析: 两者的 IO 访问量之比为: IOstandardIOflash=Θ(N2)Θ(N2d2/M)=Θ(Md2)\frac{\text{IO}_{\text{standard}}}{\text{IO}_{\text{flash}}} = \frac{\Theta(N^2)}{\Theta(N^2 d^2 / M)} = \Theta\left(\frac{M}{d^2}\right) 在 NVIDIA A100 上,每个 SM 的 Shared Memory 可配置为 M≈164 KB=82,000 FP16 elementsM \approx 164 \text{ KB} = 82,000 \text{ FP16 elements}。 当 d=64d = 64 时, d2=4,096d^2 = 4,096: Md2≈82,0004,096≈20×\frac{M}{d^2} \approx \frac{82,000}{4,096} \approx 20 \times FlashAttention 将全局显存总线上的数据搬运量,实打实地削减了整整 20 倍! 这正是它为什么能一举封神的技术原力。

4. 从 FlashAttention-1 到 FlashAttention-2 的极致工程重构

尽管 FlashAttention-1 在理论上极其完美,但在 2022 年落地大厂生产线时,系统工程师们发现它的实测性能并没有达到硬件的物理极限:在 A100 上,其算力利用率(MFU)大约停留在 30%45%(约 100140 TFLOPS),距离 A100 FP16 理论上限 312 TFLOPS 仍有巨大鸿沟。 2023 年,Tri Dao 发布了 FlashAttention-2,通过对 GPU 线程级体系结构与指令流的极致工程重构,将算力利用率一举推升到了惊人的 55%~73%(突破 220 TFLOPS)! 它究竟做对了什么?

4.1 FA-1 的三大性能缺陷深剖

  1. 非矩阵乘指令(Non-matmul FLOPs)开销过高: 在 FA-1 中,由于 Online Softmax 的中间更新公式在内层循环频繁触发,线程在每个分块步长内都要执行大量的标量除法、指数运算以及针对矩阵 OO 的广播乘法缩放(Rescaling)。在 GPU 微架构中,Tensor Core 只能执行纯粹的 GEMM 乘加,所有这些元素级的 Rescale 全都排队落在了低性能的标量 CUDA Core 或 Special Function Unit(SFU)上,导致 Tensor Core 频繁等待。
  2. 内外循环嵌套倒错引发的 HBM 往返写回: 在 FA-1 中,外层循环是 K,VK, V,内层循环是 QQ。 这意味着,针对同一个 QiQ_i 块的局部输出 OiO_i,在外层循环不同的 Kj,VjK_j, V_j 步长中都会被更新一次! 为了让不同的外层迭代能接力累加,FA-1 必须在每次内层循环结束时,把中间的 Oi,mi,ℓiO_i, m_i, \ell_i 写回到全局 HBM,在下一个 KVKV 块时再从 HBM 读出来!这直接破坏了局部性,带来了额外的全局读写。
  3. Warp 级并行度与因果掩码(Causal Mask)分配失衡: 在处理自回归模型的下三角因果掩码时,FA-1 粗暴地将对角线之外的块全部计算,导致大量的全零无效计算;同时 Warp 内部的切分导致了大量的 Shared Memory 读写 Bank Conflict。

4.2 架构颠覆一:内外循环翻转(Outer Loop on Q, Inner Loop on KV)

FlashAttention-2 最关键的架构重构,就是将内外循环彻底调换:
  • 外层循环遍历 QQ 分块;
  • 内层循环遍历 K,VK, V 分块。
这个翻转带来的工程红利是具有决定性的: 现在,我们将一个 QiQ_i 分块的工作量独占性地分配给 GPU 的一个 Thread Block(线程块)。
  • 这个 Thread Block 在其生命周期内,只需要从 HBM 加载一次 QiQ_i 块放入 Shared Memory;
  • 更重要的是,中间累加器 OiO_i 可以自始至终保存在线程的私有寄存器(Registers)中!
  • 内部循环不断把不同的 Kj,VjK_j, V_j 块拉进 Shared Memory,直接乘加到当前常驻寄存器的累加器上;
  • 在内层循环全部跑完之前,根本不需要向 HBM 写回任何中间数据!
  • 直到该线程块彻底算完所有的 K,VK, V 块,才在最后将完整的 OiO_i 一次性写入 HBM!

4.3 架构颠覆二:Softmax 动态缩放后置(Lazy Rescaling)

在 FA-1 中,每次更新累加器时,都在做除以 ℓ\ell 的归一化。 FA-2 对 Online Softmax 数学递推式做出了关键化简:解耦未归一化乘加与最终除法归一化。 在内层循环中,线程完全不维护归一化后的 OO,而是维护纯粹的未归一化加权累加矩阵: Oi(j)=Oi(j−1)⋅emi(j−1)−mi(j)+eSij−mi(j)VjO_i^{(j)} = O_i^{(j-1)} \cdot e^{m_i^{(j-1)} - m_i^{(j)}} + e^{S_{ij} - m_i^{(j)}} V_j 同时维护分母: ℓi(j)=ℓi(j−1)⋅emi(j−1)−mi(j)+∑eSij−mi(j)\ell_i^{(j)} = \ell_i^{(j-1)} \cdot e^{m_i^{(j-1)} - m_i^{(j)}} + \sum e^{S_{ij} - m_i^{(j)}} 注意看:在所有内层循环执行完毕前,完全不需要进行任何昂贵的高延迟除法指令! 直到整个序列的 K,VK, V 块全部迭代完毕,每个线程才执行一次: Oi=diag(ℓi)−1OiO_i = \text{diag}(\ell_i)^{-1} O_i 将昂贵的除法操作减少到了整个算子生命周期的最后一次,使整个内层主循环中 95% 以上的指令全部是吞吐极高的 Tensor Core GEMM!

4.4 架构颠覆三:Warp 协同切分与因果掩码调度

在 FlashAttention-2 中,Thread Block 内部由 4 或 8 个 Warp(每个 Warp 32 线程)组成。 FA-2 重构了 Warp 之间的任务分配:
  • 不再让不同的 Warp 去竞争同一个 QKTQ K^T 结果,而是将 QQ 的行数平均切分给不同的 Warp;
  • 所有 Warp 共享从 Shared Memory 读取的完整 K,VK, V 列切片,各 Warp 独立计算其负责的行输出,彻底消除了 Warp 之间的跨线程同步屏障(__syncthreads());
  • 因果掩码(Causal Mask)极速剪枝: 对于完全处于因果掩码下方的块(纯粹的无效未来 Token),调度器直接在 Grid 层面跳过执行;只有横跨对角线的边缘分块才执行细粒度掩码逻辑,使自回归训练/预填充的计算量直接净省近 50%!

5. FlashAttention-3 与 Hopper 硬件特性的深度融合

Ringi 导师解构:FlashAttention-1 到 FlashAttention-3 的演进天梯 2024 年下半年,随着 NVIDIA Hopper 架构(H100/H800)在大模型基础设施中的全面普及,Tri Dao 与 Colfax 团队推出了面向 Hopper 时代的算子皇冠之作——FlashAttention-3。 在 H100 上,FP16 峰值算力来到了惊人的 989 TFLOPS,同时引入了革命性的硬件新特性。FlashAttention-3 将算力利用率推向了难以置信的 75%~85%(实测接近 750 TFLOPS)! 它深度榨干了 Hopper 架构的三大全新物理硬件红利:

5.1 TMA(张量内存加速器):释放通用寄存器与指令发射器

在 Ampere(A100)架构上,从 HBM 拷贝数据到 Shared Memory 虽然有 cp.async 指令,但依然需要通用 CUDA Core 计算每一个线程要加载的内存地址,并将指令发射到流水线中,这极大地消耗了寄存器和指令发射槽位。 Hopper TMA(Tensor Memory Accelerator)彻底改变了游戏规则:
  • 它是一个纯硬件级的异步二维/多维张量 DMA 搬运引擎;
  • 线程只需要发射一条极简的高阶描述符指令(指定张量的全局基地址、分块尺寸 BM×BNBM \times BN 和跨度 Stride);
  • 硬件控制器在后台自动完成多维地址转换、边界处理并以最高总线效率将数据从 HBM 注入 Shared Memory;
  • 整个搬运过程完全不经过通用寄存器文件(Zero Register Footprint),为每个线程节约了 30~40 个珍贵的物理寄存器,直接将 SM 的 Occupancy(占用率)拉满!

5.2 WGMMA 原语:128 线程大张量群矩阵指令

在过去的架构中,MMA 指令是以单 Warp(32 线程)为粒度执行的(如 mma.sync.aligned.m16n8k16)。 在 Hopper 架构中,NVIDIA 推出了 WGMMA(Warp Group MMA):
  • 以 Warp Group(4 个 Warp,共 128 个连续线程)为一个不可分割的协同运算引擎;
  • 输入矩阵 BB(如 KK 或 VV 矩阵)直接存放在 Shared Memory 中,WGMMA 可以绕过寄存器直接从 Shared Memory 读取操作数,直接执行如 64×128×1664 \times 128 \times 16 规模的巨型矩阵乘!
  • 相比旧版 MMA,WGMMA 的指令发射开销骤降 4 倍,彻底释放了计算管线的发射带宽。

5.3 Warp 特化(Warp Specialization)与双缓冲乒乓流水线

为了让 TMA 搬运与 WGMMA 计算达到完美的重叠,FlashAttention-3 彻底摒弃了传统的所有线程“既当搬运工、又当计算员”的模式,采用了尖端的 Warp Specialization(Warp 特化):
  • 在一个 Thread Block(例如 160 个线程)中:
    • 1 个 Producer Warp(32 线程):专职负责与硬件 TMA 对话,负责计算全局边界并源源不断地向后台发射下一个分块的异步预取请求;
    • 4 个 Consumer Warps(128 线程,组成一个 Warp Group):纯粹作为数学计算引擎,只执行 WGMMA 和极简的 Online Softmax 更新;
  • 两者通过硬件级异步事务屏障(cuda::barrier / mbarrier)进行同步。计算引擎在翻炒当前 Tile 的同时,硬件 TMA 已经在把下一个 Tile 塞入 Shared Memory 的 Ping-Pong 缓冲区,在时间轴上实现了对 HBM 访存延迟的 100% 完美掩盖!

5.4 FP8 低精度支持与误差控制

FlashAttention-3 还是业界首个将生产级 FP8 引入超长上下文 Attention 的工业内核:
  • 采用 FP8(E4M3 / E5M2)进行 QKTQ K^T 和 PVP V 矩阵乘法,理论吞吐直接翻倍至近 2 PFLOPS;
  • 针对 Softmax 敏感问题,采用 分块动态 Scale Factor:在 Shared Memory 内部将每行的数据动态放大归一化,防止小数值在 FP8 下发生严重精度欠拟合,在保持大模型长文本大海捞针(Needle in a Haystack)100% 准确率的同时,实现了训练与推理吞吐的又一次翻倍跃迁。

6. 动手实战与代码实验室(Minimal Runnable Code)

遵循 Full-Output Enforcement 铁律:以下所有代码均为生产可编译运行的完备实验程序,绝无任何 // ... 或占位省略!

实验 1:Standard Attention vs FlashAttention 显存开销与 IO 流量微基准(Python)

本实验通过 PyTorch 原生 API 与底层内存分配器跟踪,直观量化随着序列长度 NN 增长, O(N2)O(N^2) 与 O(N)O(N) 之间的显存与耗时鸿沟。 保存为 attention_io_memory_benchmark.py 并运行:

实验 2:C++/CUDA Tiled Attention 核心原理最小实现(CUDA Kernel)

本实验用原生的 CUDA C++ 编写一个极度简化的单头 Tiled Forward Attention Kernel,完整展示 Shared Memory 分块、在线最大值追踪与累加器 Rescaling 机制。 保存为 tiled_attention_forward.cu,使用 nvcc -O3 -arch=sm_80 tiled_attention_forward.cu -o tiled_attn 编译运行:

实验 3:OpenAI Triton 工业级 FlashAttention-2 Forward Kernel 完整实现

在现代大模型工程中,OpenAI Triton 已经成为编写高性能 Attention 算子的事实工业标准。 以下代码实现了一个工业级完整的 FlashAttention-2 前向算子,完整包含:
  • 外层循环 QQ,内层循环 KVKV;
  • 惰性缩放(Lazy Rescaling);
  • 因果掩码(Causal Mask)优化跳过;
  • 支持任意 Head Dimension 的 Padding 保护。
保存为 triton_flash_attention_v2.py 并运行:

实验 4:端到端吞吐、延迟与 TFLOPS 基准压测脚本

本脚本自动化遍历不同序列长度(1K 到 32K),精确测算标准 Attention、PyTorch 官方 SDPA 与我们手写的 Triton FlashAttention-2 的真实吞吐、端到端延迟与实测 TFLOPS。 保存为 flash_attention_end_to_end_bench.py 并运行:

7. Ringi 避坑指南与生产黄金准则

7.1 大厂生产避坑指南(7 大认知盲区纠偏)

7.2 生产性能工程黄金 Checklist


8. Ringi 5 点核心速记口诀、自我检验清单与课后深度思考题

8.1 5 点押韵核心速记口诀

8.2 10 条白板自我检验清单

  1. 能否在白板上默写出 Safe Softmax 与 Online Softmax 的核心数学推导?
  2. 为什么 Online Softmax 的更新因子 emprev−mnewe^{m_{\text{prev}} - m_{\text{new}}} 永远不会发生浮点数上溢(Overflow)?
  3. 在 FlashAttention-1 中,为什么 HBM 访问复杂度是 Θ(N2d2/M)\Theta(N^2 d^2 / M),而不是 Θ(N2)\Theta(N^2)?
  4. 为什么 FlashAttention-2 将外层循环改为遍历 QQ,就能完全消除跨 Block 的同步与写回?
  5. FlashAttention-2 是如何实现 Lazy Rescaling(延迟缩放)的?它省去了哪些底层的 GPU 运算指令?
  6. 在自回归长文本训练中,因果掩码(Causal Mask)是如何在 Block 粒度实现调度剪枝的?
  7. 为什么反向传播“丢弃 PP 矩阵再重算一次”的速度,反而比“把 PP 矩阵保存在显存里直接读”要快得多?
  8. Hopper 架构中的 TMA(张量内存加速器)相比 Ampere 的 cp.async,在寄存器开销上有何本质飞跃?
  9. 为什么在大模型推理的生成(Decode)阶段,直接使用原版的 FlashAttention-2 性能反而会变差?
  10. 当 Head Dimension d=128d=128 时,为什么传统 Attention 算子的算术强度被锁死在约 85 FLOPs/Byte?

8.3 3 道大厂高阶课后深度思考题

  1. FlashDecoding 的跨 Block 树状规约开销:在 Decode 阶段,由于 QQ 只有 1 个 Token,FlashDecoding 将超长的 K,VK, V 序列切分成多个 Split 分配给不同的 SM 并行处理,各个 SM 独立算出一个局部输出 OsplitO_{\text{split}} 与局部统计量 (msplit,ℓsplit)(m_{\text{split}}, \ell_{\text{split}})。请设计一个高效的跨 SM 归并算法(如基于原子操作或两阶段规约),并分析当 Split 数量达到 128 时,规约开销与并行度收益的平衡点在哪里?
  2. PagedAttention 与 FlashAttention 的天作之合:在 vLLM 的 PagedAttention 机制中, K,VK, V Cache 在物理显存中是不连续的非物理页(Pages)。如果要将 FlashAttention-2 的 Triton 内核与 PagedAttention 融合,内核的内存地址计算与加载流水线应做出何种改动?如何避免非连续跨页读取导致的内存合并访问失效?
  3. FP8 Attention 的数值下溢(Underflow)死穴:在 FlashAttention-3 中使用 FP8(E4M3)格式存储 SS 和 PP 矩阵时,由于 E4M3 的动态范围极其有限(最小非规格化数约为 2−92^{-9} ),在经过 Softmax 减去最大值后,大量稍小的非核心 Attention 权重会直接下溢变成 0,导致长文本检索任务中细微线索丢失。请提出一种基于局部块动态缩放(Per-block Quantization Scale)或混合精度累加的工程补救方案。

9. 📚 参考资料与核心源码/经典论文指引

  1. 经典开山论文:
    • Dao, Tri, et al. “FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness.” NeurIPS 2022.
    • Dao, Tri. “FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning.” ICLR 2024.
    • Shah, Jay, et al. “FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision on Hopper GPUs.” arXiv:2407.08608, 2024.
  2. 本地核心知识库代码对照(AI_BOOK):
    • 理论与演进:vLLM FlashAttention 模块深度分析(AI-fundamentals)
    • 算法手算推导:FlashAttention V1 详解与算法演进(AIInfraGuide)
    • 算子实战实现:LeetCUDA Attention 优化源码(LeetCUDA)
  3. 官方开源项目与核心代码实现:

附录:Appendix A — 大厂硬核高频面试题与白板推导(Interview Drill)

题目 1:白板推导 Online Softmax 递推公式,并解释为什么能从两遍扫描变为单遍流式计算

考核大厂:字节跳动(字节跳动 AML / 大模型系统架构)、阿里巴巴(通义千问 Infra 组)
解题标准答案:
  1. 两遍扫描的根本原因: 标准 Softmax 的归一化公式为:
pi=exi−mℓ,m=max⁡1≤k≤Nxk,ℓ=∑k=1Nexk−mp_i = \frac{e^{x_i - m}}{\ell}, \quad m = \max_{1 \le k \le N} x_k, \quad \ell = \sum_{k=1}^N e^{x_k - m} 由于 mm 必须通过遍历全局所有 NN 个元素才能确定,因此第一遍扫描必须先求出 mm;在 mm 确定后,才能进行第二遍扫描求出每个元素的指数 exi−me^{x_i - m} 并累加得到分母 ℓ\ell。 2. Online Softmax 的流式化解耦: 假设当前处理到了第 kk 个元素,已知前 k−1k-1 个元素的最大值为 mk−1m_{k-1},局部配分和为 ℓk−1=∑j=1k−1exj−mk−1\ell_{k-1} = \sum_{j=1}^{k-1} e^{x_j - m_{k-1}}。 新读入第 kk 个元素 xkx_k,新的局部最大值为: mk=max⁡(mk−1,xk)m_k = \max(m_{k-1}, x_k) 将前 k−1k-1 项的指数基准从 mk−1m_{k-1} 平移至 mkm_k: ℓk=∑j=1kexj−mk=∑j=1k−1e(xj−mk−1)+(mk−1−mk)+exk−mk=ℓk−1⋅emk−1−mk+exk−mk\ell_k = \sum_{j=1}^k e^{x_j - m_k} = \sum_{j=1}^{k-1} e^{(x_j - m_{k-1}) + (m_{k-1} - m_k)} + e^{x_k - m_k} = \ell_{k-1} \cdot e^{m_{k-1} - m_k} + e^{x_k - m_k}
  1. 架构意义: 通过维护衰减倍率 α=emk−1−mk\alpha = e^{m_{k-1} - m_k},算法在单遍扫描过程中,无需预先知晓全局最大值,即可实时维护正确的归一化分母与累加器状态。将时间复杂度保持在 O(N)O(N) 的同时,将数据扫描次数从 2 次降为 1 次,并完全消除了对全局显存缓存中间状态的依赖。

题目 2:白板手算对比:在序列长度 N=8192N=8192、头维度 d=128d=128、SRAM 大小 M=100KBM=100\text{KB} 下,Standard Attention 与 FlashAttention 的理论 HBM 访存量各为多少?降低了多少倍?

考核大厂:腾讯(混元大模型 Infra)、商汤科技(大装置算子架构)
解题标准答案:
  1. Standard Attention 访存量手算: 每个数据以 FP16(2 字节)存储。
    • 读取 Q,K,VQ, K, V:
3×(N×d×2)=6×8192×128=6.29×106 Bytes≈6.29 MB3 \times (N \times d \times 2) = 6 \times 8192 \times 128 = 6.29 \times 10^6 \text{ Bytes} \approx 6.29 \text{ MB}
  • 写入中间矩阵 SS:
N×N×2=(8192)2×2=134.22×106 Bytes≈134.22 MBN \times N \times 2 = (8192)^2 \times 2 = 134.22 \times 10^6 \text{ Bytes} \approx 134.22 \text{ MB}
  • 读取 SS 并写入 PP(Softmax): 2×(N×N×2)≈268.44 MB2 \times (N \times N \times 2) \approx 268.44 \text{ MB};
  • 读取 PP 并写入输出 OO: N×N×2+N×d×2≈134.22 MB+2.10 MBN \times N \times 2 + N \times d \times 2 \approx 134.22 \text{ MB} + 2.10 \text{ MB};
  • 总访存量:
Totalstd≈4×(2N2)+8Nd=4×134.22 MB+8.39 MB≈545.27 MB\text{Total}_{\text{std}} \approx 4 \times (2 N^2) + 8 N d = 4 \times 134.22 \text{ MB} + 8.39 \text{ MB} \approx 545.27 \text{ MB}
  1. FlashAttention 访存量手算:
    • SRAM 大小 M=100 KB=102,400 Bytes=51,200 FP16 elementsM = 100 \text{ KB} = 102,400 \text{ Bytes} = 51,200 \text{ FP16 elements};
    • 分块大小 Bc≈M4×d=512004×128=100B_c \approx \frac{M}{4 \times d} = \frac{51200}{4 \times 128} = 100 个 Token;
    • 外层循环次数 Tc=⌈N/Bc⌉=8192/100≈82T_c = \lceil N / B_c \rceil = 8192 / 100 \approx 82 次;
    • 每次外层循环必须读取一次完整的 QQ 矩阵:
82×(N×d×2)=82×2.10 MB≈172.2 MB82 \times (N \times d \times 2) = 82 \times 2.10 \text{ MB} \approx 172.2 \text{ MB}
  • K,VK, V 仅在外层循环加载一次: 2×(N×d×2)≈4.2 MB2 \times (N \times d \times 2) \approx 4.2 \text{ MB};
  • 最终输出 OO 仅在最后写回一次: N×d×2≈2.1 MBN \times d \times 2 \approx 2.1 \text{ MB};
  • 总访存量:
Totalfa≈172.2+4.2+2.1=178.5 MB\text{Total}_{\text{fa}} \approx 172.2 + 4.2 + 2.1 = 178.5 \text{ MB}
  1. 访存削减倍数:
Reduction Ratio=545.27 MB178.5 MB≈3.05×\text{Reduction Ratio} = \frac{545.27 \text{ MB}}{178.5 \text{ MB}} \approx 3.05 \times 若序列长度增加至 N=32,768N=32,768(32K):
  • Standard 访存量随 N2N^2 放大 16 倍,膨胀至 8.72 GB;
  • FlashAttention 访存量仅增加至约 1.2 GB;
  • 访存削减倍率进一步飙升至 7.3 倍!

题目 3:为什么 FlashAttention-2 将外层循环设为 QQ、内层循环设为 KVKV 就能消除原子锁与跨 Block 同步?

考核大厂:美团(基础模型算子团队)、NVIDIA(GPU Computing 架构面试)
解题标准答案:
  1. FA-1 的冲突本质(外层 KVKV,内层 QQ ): 在外层循环中,每个迭代只持有一段 Kj,VjK_j, V_j 数据。内层循环计算所有 QiQ_i 与当前 KjK_j 的点积。 这意味着,针对同一个输出行块 OiO_i,它的完整结果是由所有的 Kj,VjK_j, V_j 共同贡献的。 如果在 Grid 级别将内层循环并行化(多个 Thread Block 同时处理不同的 QiQ_i ),那么当下一个外层循环 Kj+1,Vj+1K_{j+1}, V_{j+1} 开始时,计算出的增量必须累加到同一个全局显存位置 OiO_i 上。这会导致多个不同时间片发射的线程块竞争写入相同的全局内存行,必须通过全局内存原子操作(Atomic Add)或将未完成的中间累加值频繁写回并重新读取来维持同步。
  2. FA-2 的架构解耦(外层 QQ,内层 KVKV ): FA-2 将整个输出矩阵 OO 沿行维度切分,每个行分块 QiQ_i 被独占性地分配给唯一的一个 Thread Block(线程块)。
    • 该线程块在启动后,拥有对输出子块 OiO_i 的绝对唯一写权限;
    • 在内层循环中,该线程块依次加载所有的 K1,K2,…,KMK_1, K_2, \dots, K_M,将中间贡献持续累加在线程块私有的片上寄存器(Registers) 中;
    • 由于不存在任何其他线程块试图写入该 OiO_i 区域,因此整个内层循环全程 0 原子操作、0 全局内存同步、0 中间结果 HBM 写回!
    • 计算完毕后,单次写回全局显存,彻底消除了内存竞争与同步气泡。

题目 4:FlashAttention 反向传播为什么要重新计算注意力权重 PP?算力增加与显存节省的收支平衡点在哪里?

考核大厂:百度(文心一言高性能算子组)、MiniMax(核心 Infra 组)
解题标准答案:
  1. 反向重计算的核心机理: 在反向传播中,根据链式法则需要计算 ∇V=PT∇O\nabla V = P^T \nabla O 以及由 ∇OVT\nabla O V^T 推导出的 ∇Q\nabla Q 和 ∇K\nabla K。 标准实现选择在 Forward 阶段将整个注意力概率矩阵 P∈RN×NP \in \mathbb{R}^{N \times N} 保存到 HBM 中,Backward 直接读取。 FlashAttention 则在前向阶段彻底丢弃 PP,仅在全局内存中保存尺寸为 O(N)O(N) 的行统计量(最大值 mm 和配分对数总和 LL )。在反向传播处理到当前分块时,再次从 HBM 读取分块 QiQ_i 和 KjK_j,在 SRAM 内部重新执行一次小矩阵乘法 QiKjTQ_i K_j^T 并利用保存的 mi,Lim_i, L_i 当场复原出局部 PijP_{ij}。
  2. 算力与显存的收支平衡账本:
    • 计算成本:重新计算一次 Sij=QiKjTS_{ij} = Q_i K_j^T 和 Softmax,为整个 Attention 反向传播增加了约 2N2d2 N^2 d 的 FLOPs(占 Attention 全流程总浮点运算量的约 25%~30%);
    • 显存与访存收益:
  • 显存占用:从 O(B⋅H⋅N2)O(B \cdot H \cdot N^2) 骤降至 O(B⋅H⋅N)O(B \cdot H \cdot N)。对于 32K 序列长度,单卡节省了超 30 GB 显存,使原本根本无法训练的模型得以在有限单卡中跑起来;
  • 访存时间收益:在 HBM 上写入并再次读取一个 N×NN \times N 的 FP16 矩阵,需要 4N24 N^2 字节的访存。当算术强度低于硬件平衡点(A100 上为 153 FLOPs/Byte)时,执行这 2N2d2 N^2 d 次 Tensor Core 计算所消耗的时间,远远小于在慢速 HBM 总线上搬运 4N24 N^2 字节所消耗的时间!
  1. 收支平衡点(Break-even Point): 只要序列长度 NN 满足 N≥512N \ge 512,Attention 的计算就彻底进入 Memory-bound 状态。在此区间内,重计算所增加的纯计算时间完全被消除 HBM 搬运所节省的时间所淹没。因此,重计算不仅没有变慢,反而由于消除了内存总线拥塞,端到端反向传播速度提升了 2~3 倍!