☰
AIInfra算子优化:FlashAttention数据流与工程实践
2026/10/10 10:42:58 网站建设 项目流程

做AIInfra这几年,我越来越觉得,一个算子的好坏,不看它的标称FLOPs有多高,而看它在真实机器上的数据流有没有浪费。FlashAttention就是最典型的一个例子:它没有降低attention的数学复杂度,却靠“IO感知 + 分块 + online softmax”这套组合拳,把长序列训练的显存占用和耗时拉下来一大截,也几乎成了AIInfra从业者的必修课。这篇是AIInfra笔记的第四篇,重点聊FlashAttention的数据流优化,以及我在工程验证阶段反复踩过的坑,希望对负责算子开发、性能优化或推理部署的同学有些帮助。

1. 为什么FlashAttention的数据流会成为优化重点

1.1 标准Attention实现的HBM访问账单

先看一个最朴素的标准attention流程。假设序列长度是N,每个token的向量维度是d,那么一次前向要做:

  1. 读Q、K、V到片上计算单元;
  2. 计算S = QK^T,把S写回显存;
  3. 再从显存读S,做逐行softmax得到P,把P写回显存;
  4. 从显存读P和V,计算O = PV,写回O。

这里面最扎眼的就是S和P这两个N×N的中间矩阵。它们每出现一次,就要完整地穿过HBM。以N=4096、d=128、fp16为例,一个N×N矩阵的大小是4096×4096×2字节,也就是32MB。标准实现至少写一次S、读一次S、写一次P、读一次P,这就有128MB的HBM流量,还没算Q/K/V/O本身的读写。

把这笔账折算成时间。假设HBM带宽2TB/s,光搬运这几个中间矩阵就要64微秒左右。而QK^T加PV的计算量大约是4×N²×d = 4×4096²×128,约8.6GFLOP,即便按100TFLOPS的峰值算,也只要86微秒。也就是说,一半多的时间其实都耗在了“把半成品运回仓库再取出来”这件事上,真正在算的时间反而不到一半。序列越长,S和P就越大,浪费越严重。

我经常用做饭来打比方:标准attention像厨师每炒完一个步骤就把半成品塞回冰箱,要用的时候再从冰箱里拿出来接着做。而FlashAttention的思路,是让半成品一直待在灶台边的案板(SRAM)上,全部做完再端出去。

1.2 数据流优化的目标与总体思路

FlashAttention的核心不是减少计算量,而是减少HBM访问。它把Q、K、V切成一个个block,让QK^T、softmax、PV这三步都在片上完成,不把S和P写回HBM。因为片上SRAM的带宽比HBM高一个数量级,但容量很小,所以必须分块,让每次放进来的数据不超出SRAM容量。

这种“为访存带宽而优化”的思路,业界叫IO-aware。深度学习的很多算子其实是memory-bound,FlashAttention最开始火起来,正是因为它把传统attention的HBM访问从O(N² + Nd)量级降了下来。在足够大的N下,它能把中间矩阵的读写省掉绝大部分,这才是真正的提速来源。

理解了这个前提,后面看tiling、online softmax这些细节,就不会觉得是炫技,而是自然而然的事。

2. FlashAttention的数据流算法拆解

2.1 前向pass的tiling结构

FlashAttention前向算法并不复杂,关键是分块循环。这里给一个高度简化的伪代码,方便后面讲原理:

def flash_attn_forward(Q, K, V, B_r, B_c): # Q: (N, d), K/V: (N, d) # B_r是Q的行分块大小,B_c是K/V的行分块大小 for i in range(0, N, B_r): O_i = 0 l_i = 0 m_i = -1e30 # 实际实现建议用有限负值,下文会讲 for j in range(0, N, B_c): S_ij = Q[i:i+B_r, :] @ K[j:j+B_c, :].T # (B_r, B_c) m_ij = row_max(S_ij) # 当前块每行最大值 m_new = max(m_i, m_ij) # 更新全局行最大值 P_ij = exp(S_ij - m_new) # 用全局最大值做指数减 alpha = exp(m_i - m_new) # 旧累计O的缩放因子 O_i = O_i * alpha + P_ij @ V[j:j+B_c, :] l_i = l_i * alpha + row_sum(P_ij) m_i = m_new O_i = O_i / l_i

外层循环遍历K/V分块,内层循环处理当前Q分块。也就是说,每个Q分块会“扫”一遍所有K/V分块,不断累积输出O和归一化因子l。最后,把O_i除以l_i,得到真正的attention输出。

这个双层循环的顺序是有讲究的。把K/V放在外层,是因为在causal mask场景下可以跳过大量不参与计算的块;同时,对Q的分块量可以保持在寄存器/片上内存里,反复复用,减少K/V的重复读取。不同的实现会调整内外层顺序,但“让中间S/P不写回HBM”这个原则是不变的。

2.2 online softmax的数值原理

标准softmax需要先知道一整行的最大值,才能算exp,避免溢出。FlashAttention不允许等所有K/V块都读完再算,所以它把softmax拆成了“递推合并”的形式。

每一行我们维护三个量:当前累计最大值m_i、当前累计指数和l_i、当前累计输出O_i。每来一个新块,通过下面的规则合并:

  1. 取m_new = max(m_i, m_ij);
  2. 旧输出乘上exp(m_i - m_new),调整到以m_new为基准的指数空间;
  3. 新块P_ij = exp(S_ij - m_new)直接参与;
  4. 累积l_i也做同样的缩放;
  5. 更新m_i。

最终O_i / l_i就是标准softmax的结果。整个过程不会出现大数exp溢出,因为每个块用的都是当行当前最大值,而所有指数项都不会超过1。这保证了工程上常用的fp16/bf16也能稳定跑,而不需要像标准softmax那样把整行数据先拽到寄存器里统一归约。

这里有个很容易错的小地方:初始化m_i时,不要用负无穷。虽然数学上用负无穷没有问题,但在一些编译器或特殊数值实现里,exp(-inf - x)可能出现意外行为,而且负无穷参与max时也不够直观。我习惯用一个足够小的有限负数,比如-65504(fp16能表示的有限值)或者-1e30,并确保这个值比任何真实QK^T结果都小。这样既保持了数值稳定性,也避开了-inf在硬件上的特判路径。

2.3 反向传播的额外复杂点

FlashAttention的反向比前向麻烦得多。因为前向没有保存S和P,反向时需要重新读取Q、K、V并重算每个块的S和P,再做softmax的链式推导。

反向的核心公式要处理两个分支:一是对V的梯度,dV = P^T dO;二是对Q和K的梯度,会经过softmax的雅可比变换。具体来说,对每一行有:

dS = P ⊙ (dP - rowsum(dP ⊙ P))

其中dP来自dO和V的乘积。这个式子看着简单,但在分块递推过程中,dQ和dK都会累积多个块的贡献,而且每个块的“m_new”和“l_i”都在变,所以反向也要维护类似的缩放状态,否则梯度会差一个因子。

我在实际开发中体会最深的一点是:不要试图复用前向的online softmax逻辑去“猜”反向,反向必须在纸面上先把递推公式推清楚,再落到kernel里。很多“看起来正确但梯度总是偏大”的问题,最后查出来都是把l_i当作常量,没有把l_i的梯度纳入dS的计算,或者对旧块输出O缩放时忘了同步缩放梯度的累积。

3. 工程验证:我是怎么确认一个实现没写错的

3.1 正确性验证矩阵怎么设计

对于算子开发,第一步永远是“数值得对”,性能是后话。FlashAttention的正确性验证,我一般分成三层。

第一层,构造小规模可手算的样例。比如N=4,d=2,Q/K/V用固定整数,手推标准attention结果,再和kernel输出比较。这能快速发现明显的计算错误。

第二层,用高精度参考实现做随机对比。我会写一个完全不考虑性能的朴素attention,用fp64计算,并支持mask和scale,作为ground truth。然后随机生成Q/K/V,比较输出。常用的随机分布是均值为0、方差为1的正态分布。除了常规随机数,还要加“压力测试”:让QK^T的数值范围很大,或者让某一行query与所有key都很相似,甚至让某些token的query是完全相同的向量,检查kernel是否会出现NaN、Inf或者误差不收敛的问题。

第三层,用同一框架里的官方融合attention做交叉验证。如果floating point误差在允许范围内,基本能证明实现逻辑没有问题。

对比指标上,不要只记最大绝对误差,建议同时记录:

  • 最大绝对误差;
  • 最大相对误差;
  • 余弦相似度;
  • NaN/Inf数量。

对于fp16的kernel,我通常允许输出与fp64参考值之间最大相对误差在1e-2量级,绝对误差在1e-3量级;fp32则严格得多,最大绝对误差一般要求小于1e-5。当然,具体阈值和scale、序列长度有关,最好在验证脚本里自动统计多个case的误差分布,而不是手动看一个分支。

3.2 性能数据应该怎么测才可信

性能测量的坑比很多人想象得多。一个kernel跑得快不快,不能只看一两次计时,更不能只看理论FLOPs。我在项目里总结了一套固定流程:

  1. 预热。先跑20次左右,把GPU频率、缓存状态稳住。
  2. 正式测试。连续跑50次,取中位数或P90,而不是平均值,因为平均值容易被个别调度尖刺拉高。
  3. 同步后再计时。异步kernel如果没有同步,CPU计时会严重失真。
  4. 不同case之间清理缓存。如果连续测不同长度的序列,后一个case可能“踩”在前一个case留下的L2缓存上,导致数据虚高。可以在每个case之间故意读写一个很大的buffer,把缓存冲掉。
  5. 用性能分析工具统计kernel时间和HBM读写量,而不是只看总耗时。
  6. 同时记录理论FLOPs和实际有效算力。有效算力=总FLOPs/kernel时间,这才能看出你离硬件峰值有多远。

还有一个非常容易踩的坑:只测一个序列长度。FlashAttention在短序列上可能没有明显优势,因为分块、初始化、同步的开销没有摊薄;序列变长后,HBM访问的节省才会凸显。所以至少测2K、4K、8K、16K一组数据,看趋势有没有随序列长度拉开差距。

3.3 老计算卡适配的教训

最近有同行在讨论一套可组合内核模板,把FlashAttention适配到一款2019年前后的老计算卡上。我没有直接参与那台机器,但类似的“老卡适配”问题我遇到过好几次。这里面的典型矛盾是:新开发出来的算子默认是为新卡的大共享内存、新指令集和更高带宽设计的,直接拿到老卡上,往往第一步就编译失败。

以分块参数为例。新卡的片上内存能轻松放下128×128的块,而老卡可能64×64都紧巴巴。我们在一轮实测中做过分块扫描,结果大概是这样的:

分块大小(B_r, B_c)现象
128×128本地内存溢出,编译失败或产生大量寄存器溢出
64×128可运行,但bank conflict明显,相对耗时1.35倍
64×64相对耗时1.00倍,当前最优
64×32无bank conflict,但计算块太小,矩阵乘法算力不足,相对耗时1.22倍

这个结果说明,分块并不是越大越好,也不是越小越好。块太大,装不下或者冲突严重;块太小,计算强度不够,访存比例又上去了。最佳点取决于老卡的本地存储容量、矩阵指令宽度和访存带宽。如果只用一套“通用默认参数”,性能大概率惨不忍睹。

另外,老卡往往不支持新式张量指令里的一些快速转置和原子特性。比如QK^T需要把K做成列访问,不支持的卡只能用普通访存方式硬转,开销很大。变通办法是调整数据布局,让K/V在片上以适合当前卡的方式排列,或者把K/V的block维度与Q的block维度对调来减少转置。这类细节必须实际在硬件上profile,没法只看规格表拍脑袋。

3.4 调试常见问题速查表

工程验证阶段的大多数问题都有规律。我整理了一张速查表,基本能覆盖最常见的症状。

症状可能原因排查方向
输出全NaNm_i初始化用了-inf,或exp的输入太大把m_i初始化为有限负值;检查scale是否先乘QK^T
小序列正确,大序列错误索引溢出或分块边界判断错误检查循环边界,尤其causal skip时是否有index越界
误差在4e-3左右,且不随dtype改变而改善online softmax更新O时忘记对旧累积值做缩放核对O = O * exp(m_old - m_new)这一步
数值正确但性能低于标准实现分块过小、未做双缓冲、bank conflict严重做分块扫描;观察本地内存访问冲突;启用多级流水
causal mask结果不对遮罩处理方式错误,比如mask位置加了有限大数而不是清零对mask位置直接清零,不要参与exp和行max
反向梯度误差大反向递推时把l_i当作常量,漏掉了l_i的梯度重新推导dS时考虑d(l_i)对O和梯度的贡献

4. 一次内部验证的复盘:某长序列推理场景的优化过程

4.1 场景设定和优化前的profile

去年我们内部要支持一个长序列直接推理场景,序列长度主要在4K到8K,头维度128,使用fp16。一开始直接接标准attention实现,batch稍微调大就碰到显存不足,而且耗时明显随长度平方上涨。

我们先用性能分析工具抓了kernel时间分布。结果非常直观:标准attention总共跑了4个kernel,其中写S、读S、写P、读P占了大约60%的HBM流量。显存方面,光S和P两个矩阵就占据了绝大部分临时存储,batch=8时已经非常紧张。

这一步profile做完,结论就很清楚:不需要动模型结构,只需要把attention的中间矩阵从HBM上消掉。FlashAttention的分块策略正好匹配这个瓶颈。

4.2 分块参数调整前后的数据

我们把参考实现切成可组合kernel模板,然后先做数值对齐,再做性能扫描。最初直接抄了新卡上的128×64参数,结果在老计算卡上编译不通过。把块降到64×64之后,才跑通。

调整后的对比数据大致如下(相对时间,以优化前基线为1.00):

方案相对耗时显存占用
标准attention,无mask优化1.001.00
标准attention + 融合中间结果0.880.90
FlashAttention 64×640.720.61
FlashAttention 64×64 + causal skip0.660.61

在causal场景下,我们还做了三角形循环跳过:一旦K块对应的token位置全部晚于当前Q块,就直接进入下一个块,不执行矩阵乘法和softmax。这一步几乎不影响正确性,但把无效计算砍掉不少。

4.3 反向验证中踩过的坑

前向验证通过后,我们开始测反向梯度。第一次用数值梯度校验时,最大相对误差到了百分之几,显然无法接受。排查过程花了很长时间,最后发现是反向update时用了“近似”的softmax梯度公式,把l_i当作前向结束后的常量,没有参与dS的链式推导。

修正方式是回到数学推导:softmax的雅可比里,分子分母的依赖必须完整展开。dS的公式里要同时包含当前块的P_ij、当前行的累计l_i,以及dP与P的内积。把这个修正后,梯度误差降到了1e-4以下,这才敢往训练里接。

5. 我在AIInfra里做算子优化的几点心得

写到这里,我想把几个反复验证过的原则单独拎出来说。

第一,不要迷信理论FLOPs,要算实际搬运的字节数。大量算子是memory-bound,FlashAttention就是因为精准地削减了HBM访问而成功。分析任何算子时先问:哪些中间结果可以不落HBM?哪些循环可以重排来增加数据复用?

第二,改一个地方就跑一次回归。kernel优化很容易“拆东墙补西墙”。每次只改一个变量,性能数据记录下来,数值测试跑一遍,再决定是否保留。否则出了问题根本没法定位。

第三,跨硬件适配,先做分块参数扫描,再谈其他优化。不要拿新卡的参数硬套老卡。本地内存容量、bank冲突、矩阵指令宽度都不同,一个简单的二维参数扫描往往能出奇效。

第四,正确性是“抠”出来的。online softmax里的每一个alpha、每一个l、每一处m_new都必须能写清楚数学含义。任何“好像不影响精度”的简化,最后都会在极端case里找回来。

说实话,FlashAttention本身并不算一个特别复杂的kernel,它厉害在精准抓住了attention访存密集的本质。这个思路也能延伸到很多地方,比如带稀疏mask的注意力、多组查询头合算、或者把分页推理和分块缓存结合起来。这就是AIInfra最有趣的部分:一个数据流上的小改进,最后能撬动整个推理引擎的效率。我自己在写kernel的时候,也一直保留着先画“数据移动路径”的习惯,这比先画计算图要实用得多。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询