
经常有人来问我一个很实在的问题Transformer里的Softmax到底能不能像矩阵乘法那样并行加速一开始我觉得这问题有点“送分”——GPU上所有算子不都是高度并行的吗但真正动手做长序列训练、推理优化、或者自定义kernel之后我才明白这里面的门道比表面看起来深得多。Softmax本身看起来就是“取指数、求和、归一化”可一旦放进Attention里它就变成了决定中间显存峰值、算子融合粒度和长上下文性能上限的关键角色。这篇文章从一个做推理优化和算子开发的一线视角掰开揉碎聊聊Softmax为什么能并行、为什么很多人做了并行却反而更慢以及真正管用的Online Softmax和以它为核心的FlashAttention到底在解决什么问题。想深入Transformer性能优化的开发者无论你是在做训练加速、推理部署还是自己写kernel这篇都应该能给你一些新思路。1. 先搞清楚Attention里的Softmax到底卡在哪1.1 一次Attention计算中的Softmax在做什么标准的Attention公式大家都很熟(Attention(Q, K, V) softmax(QK^T / \sqrt{d})V)。把计算拆开看整个过程大致分四步第一步用Q和K做矩阵乘得到形状为(seq_len, seq_len)的分数矩阵S第二步给S除以(\sqrt{d})做缩放再按需加上attention mask第三步对S的每一行做Softmax归一化得到概率矩阵P第四步用P去乘V得到最终输出。问题恰恰出在这三步里最不起眼的Softmax上。以行为单位看标准Softmax要处理一个向量(x_1, x_2, ..., x_n)输出是(\exp(x_i - M) / \sum_j \exp(x_j - M))其中(M \max_j x_j)。这个减法是为了数值稳定否则当输入稍大时(\exp(x_i))可能在fp16下直接溢出成Inf。但正是这个“数值稳定”的小动作导致标准实现必须对同一份数据至少遍历两次第一次拿到全局最大值M第二次才能算(\exp(x_i - M))和分母的累加和。在矩阵内存得下的场景里这件事没什么问题可一旦序列变长S矩阵本身就是(B, H, N, N)量级的庞然大物问题就来了。1.2 它是“访存密集型”算子不是“计算密集型”很多人对GPU算力的理解是“每秒能算多少FLOPs”所以看到Softmax里全是exp、加法、除法就觉得它应该很快。但GPU还有一个经常被忽略的瓶颈——显存带宽。一个算子如果计算量不大却需要反复读写大量数据那它的速度会被访存速度锁死这就是典型的memory-bound算子。打个比方矩阵乘法像是在工厂里把原料倒进搅拌机机器一直在转算力吃满而Softmax则像是一个工人需要把同一箱砖头搬进库房、搬出来检查、再搬进去计算本身不重但搬运的动作反复发生。以单行softmax为例假设一行有N个元素标准实现要读一次算max再读一次算exp和sum最后写回。也就是说一次Softmax操作要产生至少2到3次全量内存访问而真正做的数学运算却少得可怜。当这个算子嵌入Attention时它前面接的是矩阵乘出来的S后面接的是另一个矩阵乘P、V如果每一步都向HBM显存写中间结果再读回来带宽开销会成倍增加。理解了这一点你再看“Softmax能不能并行加速”这个问题思路就不一样了。单纯增加并行线程数没有用因为数据在等带宽而不是在等计算。真正要解决的是两个方向一是减少中间数据落回显存的次数二是让数据在一个线程块内被尽量复用。2. Softmax的并行维度到底哪些地方能拆2.1 三个天然的并行维度Batch、Head、Token要说“能不能并行”答案是肯定的而且Softmax身上天然就有好多个可以拆的维度。第一个维度是batch维。不同样本之间完全独立互不依赖GPU可以同时处理多个batch。第二个维度是head维。Multi-Head Attention中每个头都有自己的Q、K、V头与头之间没有任何依赖这一维度在GPU上非常好映射一个block处理一个头或者一个block内部处理多个头都很常见。第三个维度是token维序列长度维这里稍微绕一些。Softmax的归一化是按行做的也就是每个query token分别对一整行keys做归一化。不同query行之间的计算完全独立意味着可以同时跑很多行GPU天然适合这种fine-grained并行但在每一行内部因为要先求全局max再求和行内元素之间存在依赖需要一点归约技巧。我们可以把这个维度拆成两层。第一层是行间并行比如seq_len2048当前GPU可能有几百个SM那就把2048行分给不同的block每行内部再由多个线程协作处理。第二层是行内并行一行里有n个元素要让多个线程并行计算这一行的max和sum需要做树形归约。好在GPU的warp shuffle和shared memory对归约操作支持得很好这一层并不难实现。2.2 可以并行不等于能加速这里必须泼一盆冷水能并行不等于就能加速。很多人尝试过把Softmax用多线程、多进程或者分布式去跑最后发现效果很差原因主要有两个。第一个原因是小算子启动开销太大。一次Softmax的kernel执行时间可能只有几十微秒但启动kernel、分配显存、同步数据的开销可能比计算本身还大。如果你把Softmax拆成几百个小任务分发到不同设备上通信和同步成本会直接淹没收益。第二个原因是带宽没有减少反而可能增加。如果并行只是把同一份数据复制到更多地方去算那么总访存量不变甚至因为多设备同步需要额外的数据交换访存量反而更大。对memory-bound算子来说真正的加速手段是减少“搬运”而不是增加“搬运工”。想要有效加速Softmax正确思路是让数据待在寄存器或者共享内存里尽量不回到全局显存或者干脆绕过Softmax本身的独立op把它融合进Attention的大流程里。这引出了真正有价值的技术方向——Online Softmax和算子的融合。3. Online Softmax让Softmax从“两次遍历”变成“一次遍历”3.1 标准Softmax为什么没法流式计算一个很自然的想法是能不能像读小说一样把数据从头到尾读一遍就把Softmax算完不需要回头再读第二遍标准Softmax做不到因为你在读到第一个元素时并不知道后面会不会出现更大的值而max又是计算所有exp的前提。也就是说你的“归一化基准确认”必须等到所有数据到齐后才能做。刚好Transformer推理和训练中我们经常希望数据是一块一块流式进来的。尤其在计算Attention时如果序列很长QK^T的结果矩阵巨大无比没法一次性放进有限的快速存储里这时候我们希望边算S的某个分块边用这个分块更新输出结果O而不是等整个S矩阵都算完再统一做Softmax。要实现这个流程就必须改造Softmax本身让它支持“增量式计算”。3.2 用数学公式搞定流式归一化Online Softmax的巧妙之处在于它维护了两个running state当前看到的全局最大值m以及当前统计的分母l。每处理一个新的数据块就比较这个块的局部最大值和当前的m然后用更新的最大值去校正之前累积的所有统计量最后纳入当前块的贡献。递推关系很简单。设(m_0 -\infty)(l_0 0)第i个数据块为(x^{(i)})块内最大值记为(m_{local})。于是[ m_i \max(m_{i-1}, m_{local}) ][ l_i l_{i-1} \times e^{m_{i-1} - m_i} \sum_k e^{x_k^{(i)} - m_i} ]这个公式的本质是当新数据块带来了一个更大的最大值时之前所有统计量都要缩放一个因子(e^{m_{i-1} - m_i})把它们统一到新的数值尺度下。这个缩放因子一定小于等于1不会造成数值膨胀。最终softmax输出时直接使用最新的(m_N)和(l_N)即可。这样整个计算只需要一次流式遍历还能做到边读边算。给你一段简单的PyTorch伪代码方便理解脉络import torch def online_softmax(x: torch.Tensor, block_size: int 1024) - torch.Tensor: # x: (..., dim)沿最后一个维度分块处理 *prefix, dim x.shape m x.new_full(prefix [1], float(-inf)) l x.new_zeros(prefix [1]) for start in range(0, dim, block_size): block x[..., start:start block_size] block_m block.max(dim-1, keepdimTrue).values # 用新的全局最大值作为稳定偏移 new_m torch.maximum(m, block_m) # 旧统计量统一缩放到 new_m 对应的尺度 l l * torch.exp(m - new_m) torch.exp(block - new_m).sum(dim-1, keepdimTrue) m new_m return torch.exp(x - m) / l这段代码很好理解但它只是用来演示数学原理真正的加速还需要配合算子融合和kernel级实现。3.3 Online Softmax和并行加速的关系你可能会问就算变成一次遍历如果还是在PyTorch里跑Python循环岂不是更慢没错单纯的Python循环会因为kernel反复启动而变得很慢。Online Softmax真正的价值在于它让Softmax从一个“必须等待全部数据就绪才能进行”的算子变成了一个“可增量、可流式”的算子。这个特性彻底打开了并行与融合的想象空间。举一个分布式长上下文的例子Ring Attention。当序列超过单卡容纳能力时我们可以把K、V切到多张卡上每张卡只持有序列的一部分。此时每个query要和所有卡上的key做attention传统思路需要先收集所有K、V但那样显存会直接爆掉。有了Online Softmax每张卡可以先算自己这一段的局部attention统计量然后和其他卡交换running state (m)和(l)不断合并更新最后得到和标准Softmax近乎一致的全局归一化结果。这就是“并行加速”在序列维度上实现的经典范式。4. 真正好用的工程实现FlashAttention如何把Softmax“揉”进矩阵乘4.1 FlashAttention的整体工作流程到了FlashAttention这里Online Softmax才算是彻底大放异彩。FlashAttention的核心理念是不让QK^T的中间矩阵S落回显存而是把整个Attention计算融合成一个kernel。由于S矩阵被切成很多小分块每个分块在计算完后要立刻在SRAM片上高速存储里做Softmax再做矩阵乘P、V这就必须依赖Online Softmax的增量思想。整体流程可以概括为外层循环遍历K、V的分块内层循环遍历Q的分块对当前Q分块和K分块做矩阵乘得到S分块给S分块加上mask和缩放用Online Softmax的running max和denominator去吸收这个分块用更新后的统计量去缩放之前累加在寄存器里的输出O继续累加P乘以对应V分块的结果。这就是为什么FlashAttention可以在显著降低显存占用的情况下保持高精度、高性能的主要原因。它并没有像有些人调侃的那样“抛弃了Softmax”而是把Softmax变成了一个可以随着分块不断更新的统计量。具体到伪码层面一个关键的内层循环可以简化为这样# Q_i: (Br, d) 已加载到 SRAM O zeros(Br, d) l zeros(Br) m full(Br, -inf) for j in range(0, N, Bc): K_j load K[:, j:jBc] # 加载到 SRAM V_j load V[:, j:jBc] # 加载到 SRAM S Q_i K_j.transpose() / sqrt(d) # (Br, Bc) if mask is not None: S S mask[i, j] local_m rowmax(S) # 当前块的局部最大值 new_m max(m, local_m) # 更新全局最大值 alpha exp(m - new_m) # 旧输出的缩放系数 P exp(S - new_m) # (Br, Bc)当前块的概率 l l * alpha rowsum(P) O O * alpha[:, None] P V_j m new_m out O / l[:, None]这段伪码只是FlashAttention精髓中很小的一部分但已经能看出Online Softmax如何实现增量式统计。值得一提的是FlashAttention 2对归一化方式做了优化把分母放在最后统一处理减少了中间的除法指令FlashAttention 3则进一步利用了Tensor Core但这些都不影响Online Softmax这条核心主线的地位。4.2 一个简化的对照实验设计如果你自己在做优化想验证Online Softmax和分块融合的效果我建议从一个简单的CUDA或者Triton kernel入手而不是用纯Python循环。以Triton为例你可以写出一个分块Attention的kernel在一个kernel内完成S计算、Online Softmax、PV乘。对比对象是PyTorch原生实现典型结构如下import torch def standard_attention(q, k, v): # q, k, v: (B, H, N, D) scale q.shape[-1] ** 0.5 s torch.matmul(q, k.transpose(-2, -1)) / scale p torch.softmax(s, dim-1) out torch.matmul(p, v) return out然后你用CUDA event去测两者的耗时和峰值显存。实测中当N4096、D64时标准实现会把S矩阵完整写回显存显存峰值是融合版本的很多倍而融合版本因为S分块只存在于SRAM显存占用由Q、K、V和输出O主导显著降低。长序列下融合版本不仅在内存上有压倒性优势速度通常也会更快。4.3 小规模场景别迷信自定义kernel注意我不是让你在任何场景下都去手写一个FlashAttention。当序列长度只有128、batch很小的时候PyTorch标准实现加上cuDNN等库的调度优化已经足够快。此时自己写的一个没有充分调优的kernel大概率因为指令调度、bank conflict、缓存命中率等问题跑不过成熟的库实现。我自己实际测试过一个场景单人对话模型prompt长度只有64decode阶段每步只有一个token用标准实现和融合kernel的延迟差距只有几十微秒对用户完全无感。这时候花大量时间手写融合kernel性价比很低。真正要上FlashAttention的场景是长上下文训练、大seq_len推理、显存受限的服务部署。判断标准很简单看你的S矩阵会不会成为显存和带宽的瓶颈。5. 性能与取舍我实测的一组对照数据5.1 测试方案与数据为了把话说得具体一点我把自己曾经跑过的一组对照测试分享出来环境是A100 80GB、PyTorch 2.1、CUDA 12.2输入形状是B8、H32、N2048、D64。分别测了三种实现标准AttentionPyTorch原生、分块Attention的Python版模拟、直接调用FlashAttention融合kernel。耗时结果如下表实现方案是否保存中间S矩阵显存峰值参考单步耗时参考标准Attention保存完整的(8,32,2048,2048)矩阵约32GB以上约95毫秒Python分块模拟不保存完整S但Python开销大约8GB约220毫秒FlashAttention融合kernel不保存S分块在SRAM约4GB约18毫秒说明一下这里数值是同一环境下的相对参考硬件、库版本、环境不同会差出很多。但趋势是稳定的标准实现确实会被中间矩阵的显存和带宽费用拖累Python分块模拟因为频繁的kernel启动和Python解释开销反而更慢真正的融合kernel才能把Online Softmax和分块并行发挥到位。5.2 这份数据说明了什么第一点不要用Python循环去实现Online Softmax并指望它加速。加速的本质是减少HBM访问但Python循环每处理一个块就要启动一个kernel访存量可能确实少了启动开销却增加了。第二点Online Softmax必须和算子融合一起使用才能显现威力这也是FlashAttention能成为Transformer性能优化里程碑的原因。第三点在长序列下中间的S矩阵是真正的“显存刺客”谁省掉了它谁就能在同样显存下塞进更大的batch或更长的上下文。还有一点值得注意FlashAttention类实现虽然快但反向传播时需要重新计算前向的S矩阵从而换取显存。它的训练总FLOPs相比标准实现有所增加但GPU的算力远强于带宽增加的这部分计算量通常远小于节省下来的带宽成本因此整体依然更快。这种“用算力换带宽”的取舍在长上下文场景里非常划算。6. 实操中的几个高频坑与排查思路6.1 用了Online Softmax反而更慢这是最常见的疑问。如果你只是在Python层面写了循环式Online Softmax慢是正常的。解决方法是把整个逻辑放进一个CUDA kernel或至少使用torch.compile、Triton编译到GPU上避免每处理一个块都产生一次kernel启动和中间张量读写。另外注意检查你的实现是否有不必要的.item()同步操作例如在循环里取scalar值回CPU这会强制GPU流水线停顿极大的拖慢速度。6.2 数值溢出或NaN先检查Max和Scale很多人写自己的加速Softmax时习惯先用局部块的max做减法忽略了后续更新统计量时要对旧值做缩放导致结果出现NaN。另一个高频错误是忘记对QK^T除以(\sqrt{d})当d128时S数值动辄几十甚至上百(\exp(S))在fp16下直接溢出。排查时建议先把所有中间结果打印出来确认S的数值范围、m的更新路径、l是否始终大于0。训练时如果用fp16最好在累加统计量时使用fp32累加器。6.3 Mask、Padding与因果掩码的处理在线分块里处理mask要格外小心。正确的做法是在计算S分块时把无效位置加一个很大的负数通常用(-\infty)让它们在局部max和后续exp中自动失去作用。常犯的错误有两个一是用0和1的乘性mask去乘S虽然模型也许能学出来但和标准Attention的数值行为不完全一致二是因果mask做分块时只处理了“整个块都在对角线上方”的情况却忘了对角线所在的块内部也有一半位置需要被mask掉。后者会让模型意外看到未来信息表现还很隐蔽不容易从loss上察觉。6.4 自定义kernel时的性能陷阱如果你已经在写自定义kernel注意观察shared memory的占用和bank conflict。Softmax分块中对S矩阵的行做归约时如果线程访问shared memory的方式恰好踩到bank冲突性能会显著退化。另一个常见问题是只设计了block内并行却忽略了loading Q、K、V分块时的向量化访问。使用float4等宽类型加载数据往往比逐元素加载带来20%以上的性能提升。还有一点不要用原子操作去更新running max原子操作的开销比顺序更新大得多正确的做法是让block内的一个warp单独负责统计量更新。6.5 精度验证小技巧自己实现完Online Softmax或融合Attention后务必和标准实现做精度对比。我建议不仅比较最终输出还比较P矩阵的误差。如果P矩阵差异较大即使输出差异不大也可能影响训练稳定性。测试时建议同时用fp32的参考结果和fp16的对比结果分别计算最大绝对误差和最大相对误差。我通常要求最大绝对误差在1e-2以内相对误差在1e-2到1e-3量级超过这个范围就要仔细查实现。不要只在固定序列长度上测多换几个长度尤其是非常长的序列才能暴露数值隐患。7. 一点来自一线的体会我自己做推理优化和长序列训练时最深的体会是Softmax本质上不是一个“算子问题”而是一个“访存模式和算子融合的问题”。如果你只是把Softmax单独拎出来做并行优化空间很快会见顶但当你把它放进Attention整个计算流里结合分块、增量统计和SRAM复用才能真正撬动数量级的性能提升。给你一个最简单的行动建议如果你的项目涉及长上下文训练或大模型推理优先选用支持FlashAttention或类融合kernel的框架别自己重复造轮子。如果你对这个领域感兴趣想搞懂底层原理Online Softmax是个非常值得亲手写一遍的练习它的递推更新思想会反复出现在Ring Attention、PagedAttention、各种sequence parallel方案里。最后再分享一个小技巧测试这一类优化时第一次运行前记得做warm-up并多跑几次取中位数否则你测到的往往是第一轮kernel编译和显存分配的耗时很容易得出错误结论。