大模型Attention算子实现与优化深度解读
大模型Attention算子实现与优化深度解读
这是一份关于大模型Attention算子在GPU上落地优化的技术笔记,内容来源于一次内部分享。从Online Softmax到Flash Attention,再到指令级优化和稀疏化方案,一条线串下来。
1. 从Softmax说起
Attention的起点永远是Softmax。公式谁都会写:
Softmax(xi)=exi∑j=1Nexj\text{Softmax}(x_i) = \frac{e^{x_i}}{\sum_{j=1}^{N} e^{x_j}}Softmax(xi)=∑j=1Nexjexi
但实际落地的时候,这三个问题躲不掉:
- 数值不稳定:xix_ixi 一大,exie^{x_i}exi 直接溢出,GPU上float16更是重灾区。
- 两次遍历:先找最大值(稳定数值),再算概率,访存压力翻倍。
- 中间结果全存:序列长度一大,NNN 个中间值全占着显存不放。
Online Softmax
所以有了Online Softmax——核心思路是用增量更新代替全量计算。每来一个新数据点,只更新两个统计量:
- mjm_jmj:当前见过的最大值
- djd_jdj:分母的指数累加和
推导过程(不展开,说结论):
dj+1=dj⋅emj−mj+1+exj+1−mj+1d_{j+1} = d_j \cdot e^{m_j - m_{j+1}} + e^{x_{j+1} - m_{j+1}}dj+1=dj⋅emj−mj+1+exj+1−mj+1
其中 mj+1=max(mj,xj+1)m_{j+1} = \max(m_j, x_{j+1})mj+1=max(mj,xj+1)。
这个变换的意义在于:不需要回看已处理的数据,一次遍历搞定,而且始终数值稳定。这是后面Flash Attention里softmax分块计算的基石。
2. GEMM:一切矩阵计算的原语
GEMM(通用矩阵乘法)的优化思路总结起来就一句话:把大矩阵切成小块,小块塞进缓存,一个线程块算一块,一个线程算一个Tile。这套逻辑在Attention里同样适用——Q × K^T 是GEMM,Score × V 也是GEMM。区别在于中间夹了一个Softmax,把两个GEMM的连续性打断了。
3. Flash Attention 核心设计
Flash Attention 解决的核心矛盾:显存太慢,缓存不够大,但Attention天生要存中间结果。
简单说就是:别把整个注意力矩阵 softmax(QKT/d)\text{softmax}(QK^T/\sqrt{d})softmax(QKT/d) 全写回显存。Q、K、V 都切成小块,每次只加载一小块到片上缓存,算完Softmax直接乘V,中间结果不落盘。这就是"Flash"的含义——快在减少了显存往返。
4. Chain GEMM:当两个矩阵乘需要连起来
Flash Attention里有两个GEMM:
Gemm0: S = Q × K^T
Gemm1: O = P × V (P = softmax(S))
问题出在 Tensor Core 的输出 Layout 和下一个 GEMM 的输入 Layout 不匹配。GPU上Tensor Core计算 mmac_f32_16x16x16f16 出来的数据分布在各个线程里,直接传给下一个矩阵乘会因为数据排列不对无法计算。
三种解决方案:
最终落地选了方案三(Gemm0多次mmac融合),具体做了三步优化:
- 调整线程持有数据:Gemm0在N方向连续四次mmac,用16x64x16的滑块,得到连续输出,省掉线程间数据交换。
- 数据加载:K方向用kpack=2,即16x64x32的滑块,一次读更多数据,向量化存取,减少指令条数。
- 减少索引计算:不需要线程间数据调整,减少了用于索引计算的向量寄存器占用。
最终整体性能提升 18%。
5. 四级Buffer流水线
这是整份材料里最精彩的部分。
核心思路四点:
- Q数据驻留寄存器:Q一次加载,在整个滑块计算完成前不释放,避免重复读显存。
- LDS拆成Tile:共享内存按一次mmac计算所需数据量切成4个Tile(编号0~3),四级流水就用四个槽位。
- 沿K维度切分GEMM:每次加载一个mmac计算的数据到对应LDS槽位,循环使用0号到3号空间。
- 异步加载 + 多GEMM流水:通过内嵌汇编控制数据加载异步执行,两个GEMM沿K维度切分后连起来用四级buffer——加载和计算相互遮掩,LDS用量减半,kernel并行度翻倍。
硬件约束下的寄存器预算:假设单个thread最多256个寄存器(1KB),block size 256,则一个block占256KB寄存器。如果芯片上总共768KB寄存器文件,occupancy = 768/256 = 3。LDS只要不超过64/3 ≈ 21KB即可——刚好够4个buffer均分。
6. Causal Mask 下的负载均衡
Causal Attention(因果注意力,即只看前面的token)有一个天然的不均衡问题:
解决办法很巧妙:让Block i 和 Block (num_blocks - i - 1) 合并到一个线程块。
同时调整M方向上的warp堆叠,尽可能调大M方向滑块,一个warp在seqlen方向算16。
7. 反向传播优化
反向的优化和正向是同构的——核心仍然是分块计算+中间结果不落盘这套逻辑。Flash Attention的反向同样避免了完整注意力矩阵的显存写入,把梯度计算也拆成了分块的在片操作。
8. FlashMLA:往前再走一步
FlashMLA是在Flash Attention基础上的进一步工程优化,几个关键点:
| 方向 | 做法 | 效果 |
|---|---|---|
| Prefill优化 | 切分QK/V,计算与数据读取overlap | Prefill持平A800 |
| Decoding优化 | 多级流水 + KV复用 | Decoding性能超越A800 |
| 量化 | 开发FP8 e5m2版本 | 相比FP16提升约25% |
| 指令重排 | 基于数据流分析优化指令顺序 | 进一步提升效率 |
| 硬件适配 | LDS优化 + 寄存器分配调优 | 降低硬件压力 |
这里的"持平A800"和"超越A800"是在特定场景下的对比,足以说明工程优化的价值——软件层面的极致优化,可以在一定条件下缩小甚至抹平硬件代差。
9. 量化与稀疏化:两条并行的优化线
9.1 量化:走FP8路线
量化选型上选择了FP8而非INT8,理由是:FP32到INT8的转换指令太多,而FP8有原生的Tensor Core高算力支持。实测FP8 e5m2版本的性能比FP16提升约25%,这个收益已经相当可观。
9.2 稀疏化:三个流派
BLASST 的核心逻辑:在块级处理中维护全局"运行最大值",每处理一个块,先算块内局部最大值,和全局最大值比较——差值小于阈值就直接跳过这个块(省掉Softmax的指数计算、Value加载和矩阵乘法)。这个过滤过程嵌入在正常计算流水中,零额外开销。
SpargeAttention 的做法不同:先压缩。高相似度的Token块合并成代表性Token,在小尺寸的近似注意力图上判断哪些区域重要,生成稀疏掩码,GPU只计算关键部分。
SLA(稀疏-线性注意力) 则是混合路线。对Q和K先做池化下采样,得到压缩注意力矩阵,然后根据权重动态分三类:
- Critical(关键):权重最大的 Top kh%k_h\%kh%,精细计算
- Marginal(边缘):中间部分,降精度计算
- Negligible(可忽略):权重最小的 Bottom kl%k_l\%kl%,直接跳过
SLA的卖点在于可微调,不是一刀切,两类方法各取所长,在质量和效率之间找了一个可控的平衡点。
10. Sparse Warp Online Softmax
这算是BLASST在Warp级别的落地实现。利用Online Softmax的增量更新特性,在GPU Warp级别动态识别并跳过远小于全局最大值的元素——这些元素对最终Softmax结果贡献为零(或接近零),省掉它们的计算就是纯赚。
三个关键操作:
- 池化:对Q和K矩阵下采样,拿到压缩的注意力权重矩阵
- 分类:根据压缩矩阵的值把原始权重块分成三类(关键/边缘/可忽略)
- 跳过:可忽略块直接不计算
整个过程嵌入在Online Softmax的常规计算流水中,不增加额外步骤,确实是"免费"的优化。
11. 总结
这份材料串起了Attention算子从算法到指令级优化的完整链路:
Softmax → Online Softmax → Flash Attention → Chain GEMM → 四级Buffer →
Causal均衡 → 反向优化 → FlashMLA → 量化(FP8) + 稀疏化(BLASST/Sparge/SLA)
几个值得记住的点:
- Online Softmax是分块计算的理论基础,没有增量更新就没有后续的一切。
- Chain GEMM的落地关键是Gemm0多次mmac融合,省掉了线程间数据交换,换来18%的性能提升。
- 四级Buffer是多级流水的典范,把21KB的LDS用到了极致,加载和计算完全遮掩。
- 量化选FP8不选INT8,看的是转换成本——FP32→INT8的转换太贵,FP8有原生Tensor Core撑腰。
- 稀疏化三条路线(动态阈值/无训练压缩/可微调混合)解决同一个问题,但切入点完全不同,值得对比着看。
原文中部分涉及内部项目名称和人员信息已做脱敏处理,技术数据和优化思路保持原样。
更多推荐
所有评论(0)