GPU 上 Softmax 算子绝大多数场景属于内存受限算子(memory‑bound),算力单元经常处于空闲状态,性能瓶颈来自 HBM 高带宽显存与片上 SRAM 之间反复读写搬运数据GitHub。Online Safe Softmax 出自 NVIDIA 2018 论文《Online normalizer calculation for softmax》,它在保证数值稳定性前提下,减少一次全局内存读取,理论较传统 Safe Softmax 取得 1.33 倍性能增益;更为重要的是,该算法支持流式分块输入,不需要拿到全部向量元素,成为 FlashAttention 能够实现分块 tiling 注意力计算不可或缺的数学基础ar5iv。很多开发者只知道 FlashAttention 省显存,却忽略其背后在线归一化算法的约束与工程代价。本文将从痛点、实现步骤、量化对比表格、工程坑点完整拆解这套算法。

一、真实工程痛点:传统 Softmax 实现的三类棘手问题

在大模型训练、推理算子开发工作中,朴素 Softmax、Safe Softmax 会遇到三类高频真实痛点,很多线上 NaN、性能不达标问题根源就在这里。

  1. 朴素 Softmax 存在数值溢出风险:当输入 logit 数值偏大,\(\exp(x_i)\)直接溢出得到 inf,除法后产出 NaN,FP16 场景尤其严重,直接导致训练中断、推理输出乱码。
  2. Safe Softmax 为保证稳定性引入额外内存访问开销:标准 Safe Softmax 必须先完整遍历一遍输入向量,算出全局最大值 max (x),再第二次遍历计算指数,第三次遍历完成归一化输出。每个向量元素对应 3 次内存读取、1 次写入,内存访问量上涨直接压垮带宽,长序列场景性能损失被放大。
  3. 传统实现不支持分块流式输入,限制大内存优化:不管朴素还是 Safe Softmax,都要求完整向量全部载入显存。当序列长度巨大,完整向量无法全部放进片上 SRAM,FlashAttention 这类分块计算框架就无法直接调用标准 Softmax,必须要有可以边读边算的在线版本。

实操观察 1:很多同学在 PyTorch eager 模式下直接调用torch.softmax,内部已经封装 Safe Softmax 逻辑。但 eager 执行会生成多个中间张量,实际内存读写次数会比理论值更高,长序列推理时 profiler 经常观测到大量 HBM 读写耗时,而 GPU SM 算力利用率很低。 实操观察 2:不少开发者做算子优化时,只盯着 FLOPs 浮点运算数量做优化,但是 Softmax 属于内存受限算子,浮点计算量很小,减少内存访问远比压缩计算量收益更大,这也是 Online Safe Softmax 的核心设计出发点。

二、三类 Softmax 实现原理与分步计算逻辑

2.1 原始朴素 Softmax

公式:

\(y_i=\frac{\exp(x_i)}{\sum_{j=1}^n\exp(x_j)}\) 计算流程:

  1. 第一次读取全部输入\(x_i\),计算指数,累加得到分母总和;
  2. 第二次读取全部输入\(x_i\),再次计算指数;
  3. 使用分母做除法输出\(y_i\)。

每个元素:2 次读、1 次写。致命缺陷:输入数值偏大时指数溢出,产生 inf/NaN,工程几乎不会直接使用

2.2 Safe Softmax(传统安全 Softmax)

为解决溢出,数学等价变换,每个元素减去全局最大值:

\(y_i=\frac{\exp(x_i-\max(x))}{\sum_{j=1}^n\exp(x_j-\max(x))}\) 此时\(\exp(x_i-\max(x))\)值域被约束\([0,1]\),不会溢出。 计算流程:

  1. 完整遍历输入,读取全部元素,求出全局 max (x);
  2. 再次读取输入,计算\(\exp(x_i-\max(x))\),累加求和分母;
  3. 第三次读取输入,计算分子,除以分母输出结果。

每个元素:3 次读、1 次写。优点是数值稳定;缺点是多一轮完整遍历,内存访问增加,必须等待全部输入加载完毕拿到全局 max,不能分块流式处理NVIDI…。

2.3 Online Safe Softmax 在线安全 Softmax

核心创新:不需要预先拿到全局最大值,在流式遍历输入的过程中,迭代维护两个运行变量:当前迭代最大值\(m_j\),归一化累积和\(d_j\)。每当发现新的更大输入值,使用指数缩放系数修正历史累积的归一化和,保证数学结果与 Safe Softmax 完全等价ar5iv。

变量定义:

  • \(m_j\):处理前j个输入后的运行最大值
  • \(d_j\):经过缩放修正后的累积归一化分母

迭代规则: 初始化 \(m_0=-\infty,\ d_0=0\) 对第\(j+1\)个输入\(x_{j+1}\):

  1. 如果\(x_{j+1}\le m_j\):最大值保持不变,直接累加新项

\(m_{j+1}=m_j,\quad d_{j+1}=d_j+\exp(x_{j+1}-m_j)\) 2. 如果\(x_{j+1}> m_j\):更新全局最大值,旧累积和需要乘以修正系数\(\exp(m_j-m_{j+1})\)做尺度对齐,之后新增当前项

\(m_{j+1}=x_{j+1},\quad d_{j+1}=d_j\cdot\exp(m_j-m_{j+1})+1\)

遍历全部输入完成后,拿到最终全局最大值\(m_V\)与总归一化因子\(d_V\),再重新遍历一次输入,计算输出\(y_i=\frac{\exp(x_i-m_V)}{d_V}\)。

内存访问统计:每个向量元素 2 次读取,1 次写入。既保留数值稳定性,又把内存读次数从 3 降回 2,理论性能达到 Safe Softmax 的\(\frac43=1.33\times\)。

实操观察 3:算法分为两轮循环,第一轮流式迭代更新\(m、d\)统计量;第二轮才生成输出。虽然减少内存读取,但每一个新元素到来时会额外执行 max、exp、乘法运算,这些额外 ALU 计算开销很小,GPU 上相比内存搬运可以忽略不计arXiv。

三、三类 Softmax 关键指标对比表

表格

实现方案数值稳定性单元素内存读写是否支持流式分块输入理论加速对比 Safe‑Softmax典型适用场景
原始朴素 Softmax差,易溢出 inf2 读 1 写支持1.33×教学演示,工程禁用
传统 Safe Softmax优秀3 读 1 写不支持,需要完整向量1.0×标准 GPU 推理 / 训练,输入完整张量
Online Safe Softmax优秀2 读 1 写✅支持流式分块1.33×FlashAttention、大词汇量 logit 计算、算子融合场景

补充说明:表格为理论内存访问统计,真实 CUDA Kernel 中会受缓存命中率、tile 分块大小、张量排布影响,实测加速不一定严格等于 1.33 倍。论文基准测试在部分向量长度下测出最高 1.3 倍加速,和理论值接近。

四、Online Safe Softmax 拓展:算子融合 Online‑Safe‑Softmax+TopK

Online Safe Softmax 遍历输入阶段已经完成全部内存读取,不需要再次从 HBM 搬运数据,因此可以和 TopK、argmax、argmin 等下游算子做内核融合(kernel fusion)。

  • 分开执行 Softmax+TopK:Softmax 写回 HBM 输出概率,TopK 再把概率重新读入,多出一轮完整读写。
  • 融合实现:在线遍历输入的时候一边维护\(m、d\),一边维护 Top‑K 候选集合;完成第一轮统计量迭代之后,第二轮直接同时输出 softmax 结果与 Top‑K 索引 / 数值,没有额外显存读写开销。

原始论文实测:Softmax+TopK 融合之后,最高可达5 倍性能提升,这个巨大收益绝大部分来自消除中间张量的 HBM 读写,而不是算法本身的计算量减少。在大模型采样阶段,logits 后接 top‑k 过滤场景,融合 Kernel 会带来非常可观的推理吞吐收益。

五、工程落地不可忽略的坑(非公开常见实操细节)

  1. 修正系数\(\exp(m_{old}-m_{new})\)的下溢风险:当新最大值远大于旧最大值,\(m_{old}-m_{new}\)是很大负数,exp 计算结果趋近 0,浮点精度丢失。FP16 环境该问题更明显,极端情况会造成分母失真。工程解决方案:可以部分中间统计量切换 FP32 保存,避免低精度下溢;或者在 log 空间维护累积和统计量,牺牲少量计算换取精度稳定性。
  2. 不是所有场景 Online Softmax 都更快:向量很短的时候,循环迭代、exp 修正带来的指令开销会盖过内存节省收益。短向量场景传统 Safe Softmax 反而性能更优,算子库需要做向量长度分支判断做动态选择。
  3. 反向传播实现难度上升:Online Softmax 前向过程没有保存完整 exp 中间结果,反向求梯度不能直接复用标准 Safe Softmax 的求导逻辑。FlashAttention 中通过重计算,在反向阶段重新跑一遍 online 迭代逻辑换取显存节省,属于以少量计算开销换取内存带宽的经典取舍。
  4. 融合算子调试难度高:Softmax+TopK 融合内核输出两份结果,调试定位精度 bug 的时候,需要分别和分开执行版本做数值对齐校验,微小浮点误差累积容易被忽略。

六、总结与落地建议

Online Safe Softmax 最大的贡献有两层:第一层是传统 Softmax 算子层面,在保证数值稳定前提下降低内存访问,取得最高 1.3 倍加速;第二层更有历史意义,它提供一套不需要完整向量、边流式输入边维护全局统计量的数学框架,直接成为 FlashAttention 分块注意力计算的基石,使得超长上下文 Transformer 训练推理成为现实。

在实际项目落地中给出三条可落地建议:

  1. 开发自定义 CUDA 算子时,长向量场景优先考虑 Online Safe Softmax;短向量直接使用传统 Safe Softmax,避免迭代修正带来额外指令开销。
  2. 如果业务链路 Softmax 之后紧跟 TopK/argmax,优先做 kernel 融合,消除中间张量 HBM 读写,这一步带来性能收益往往比单纯优化 Softmax 本身更大。
  3. FP16 推理环境,要重点监控 online 迭代过程中 exp 缩放项的下溢问题,关键统计变量建议使用 FP32 存储,防止出现隐蔽的精度退化、输出 NaN。

效率龙虾 会带着下面这段开聊

按文章《Online‑Safe‑Softmax 深度解析:从内存瓶颈到 FlashA…》把卡点收成可执行步骤:先做什么、别踩哪条、怎么验证。

用效率龙虾试这篇

本文侧重全链路风控方法论。落地时请用自身业务单据做回放验证,不要把示例阈值直接当生产策略。 相关:风控体检 · 方案资源

常见问题 FAQ

为什么说 Online Safe Softmax 是 FlashAttention 的数学基础?

因为 FlashAttention 需要分块 tiling 来节省显存,而传统 Safe Softmax 必须一次性加载整个向量才能算出全局最大值,无法分块。Online Safe Softmax 支持流式输入,可以边读边维护运行中的最大值和归一化因子,这样每个注意力块就能独立计算,无需等待整个序列数据,从而实现了高效的分块注意力机制。

Online Safe Softmax 相比传统 Safe Softmax 的理论性能提升有多大?原因是什么?

理论加速比约为 倍(即 4/3)。这主要源于内存访问优化:传统 Safe Softmax 需要三次读取输入向量(求最大值、求指数和、求输出),而 Online Safe Softmax 通过巧妙的两轮迭代(一轮维护统计量,一轮输出),将每个元素的内存读取次数从 3 次降为 2 次,显著减少了对 GPU 高带宽显存(HBM)的访问压力。

除了作为 FlashAttention 的基础,Online Safe Softmax 还能直接用于哪些算子融合场景?

非常适合与 TopK、Argmax 等需要扫描全向量的算子融合。因为其第一轮遍历已经读取并计算了所有输入,若在此过程中同时维护 TopK 候选列表,第二轮就能直接输出 softmax 结果和 TopK 索引,避免了将中间结果写回显存再读取的开销。文章提到,这种融合在大模型采样阶段能带来最高 5 倍的性能提升。

实现 Online Safe Softmax 时,有哪些工程上的常见陷阱?

主要陷阱包括:1. 数值下溢风险:当新最大值远大于旧值时,修正系数 exp(m_old – m_new) 可能下溢为零,需要使用更数值稳定的实现或处理。2. 分块大小(tile size)的选择:分块太小会降低并行度,太大则可能超出 SRAM 容量,需要根据具体硬件精心调优。3. 注意力掩码(Mask)的融合处理,这会增加 kernel 的复杂度。

在什么情况下应该选择 Online Safe Softmax 而不是传统的 Safe Softmax?

主要看两个条件:1. 内存瓶颈:如果你的场景(如长序列处理)受限于 HBM 带宽,而非计算单元,那么 Online 方案更有优势。2. 数据处理模式:如果输入可以(或必须)以流式或分块方式到达(如 FlashAttention 的 tiling 计算),那么 Online Safe Softmax 是必需的。如果数据是完整的小张量且内存充足,传统 Safe Softmax 实现更简单直接。

什么是AI智能系统?

「AI智能系统」可概括为:GPU 上 Softmax 算子绝大多数场景属于内存受限算子(memory‑bound),算力单元经常处于空闲状态,性能瓶颈来自 HBM 高带宽显存与片上 SRAM 之间反复读写搬运数据GitHub。Online Safe Softmax 出自 NVIDIA 2018 论文《Online normalizer calculation for softmax》,它在保证数值稳定性前提下,减少一次全局内存读取,理论较传统 Safe Softmax 取得 1.33 倍性能增益;更为重要的是,该算法支持流式分块输入,不需要拿到全部向量…