大模型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=djemjmj+1+exj+1mj+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:一切矩阵计算的原语

输出矩阵 C 的分块策略

输出矩阵 C

Block 0 (线程块0)

Block 1 (线程块1)

Block N (线程块N)

Thread 0: Tile(0,0)

Thread 1: Tile(0,1)

Thread K: Tile(i,j)

每个 Thread 计算一个小 Tile

GEMM(通用矩阵乘法)的优化思路总结起来就一句话:把大矩阵切成小块,小块塞进缓存,一个线程块算一块,一个线程算一个Tile。这套逻辑在Attention里同样适用——Q × K^T 是GEMM,Score × V 也是GEMM。区别在于中间夹了一个Softmax,把两个GEMM的连续性打断了。


3. Flash Attention 核心设计

Flash Attention 解决的核心矛盾:显存太慢,缓存不够大,但Attention天生要存中间结果

Flash Attention策略

存储层级

逐块加载

逐块加载

逐块写回

HBM / 显存
GB级别,带宽窄

SRAM / 片上缓存
KB~MB级别,带宽宽

Q 分块加载到 SRAM

K, V 分块加载到 SRAM

在 SRAM 内完成 Softmax + MatMul

结果写回 HBM

简单说就是:别把整个注意力矩阵 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 出来的数据分布在各个线程里,直接传给下一个矩阵乘会因为数据排列不对无法计算。

三种解决方案:

Chain GEMM 数据重排问题

方案一:Shfl

方案二:修改 Gemm1 B矩阵 Layout

方案三:Gemm0 多次 mmac 融合

三次 shfl 指令
将 mmac 结果合并连续

调整 B 矩阵输入 Layout
直接从源头适配

Gemm0 在 N 方向连续四次 mmac
16x64x16 滑块,输出天然连续

可行但开销大

可行但灵活性受限

最优方案

最终落地选了方案三(Gemm0多次mmac融合),具体做了三步优化:

  1. 调整线程持有数据:Gemm0在N方向连续四次mmac,用16x64x16的滑块,得到连续输出,省掉线程间数据交换。
  2. 数据加载:K方向用kpack=2,即16x64x32的滑块,一次读更多数据,向量化存取,减少指令条数。
  3. 减少索引计算:不需要线程间数据调整,减少了用于索引计算的向量寄存器占用。

最终整体性能提升 18%


5. 四级Buffer流水线

这是整份材料里最精彩的部分。

Gemm1 流水

Gemm0 流水

共享内存 LDS 划分 (~21KB)

Softmax

Buffer 0

Buffer 1

Buffer 2

Buffer 3

异步加载 K,V Tile → Buffer i

Tensor Core 计算 S = Q×K^T

异步加载 V Tile → Buffer i

Tensor Core 计算 O = P×V

核心思路四点:

  • 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)有一个天然的不均衡问题:

不均衡问题

大部分idle

Block 0
只需算1个位置

Block 1
算2个位置

Block 2
算3个位置

Block N
算N个位置

严重负载不均衡

解决办法很巧妙:让Block i 和 Block (num_blocks - i - 1) 合并到一个线程块

均衡后

Block 0
算位置0 + 位置N-1

Block 1
算位置1 + 位置N-2

Block k
算位置k + 位置N-k-1

每个Block的计算量相近

同时调整M方向上的warp堆叠,尽可能调大M方向滑块,一个warp在seqlen方向算16。


7. 反向传播优化

反向的优化和正向是同构的——核心仍然是分块计算+中间结果不落盘这套逻辑。Flash Attention的反向同样避免了完整注意力矩阵的显存写入,把梯度计算也拆成了分块的在片操作。


8. FlashMLA:往前再走一步

FlashMLA是在Flash Attention基础上的进一步工程优化,几个关键点:

方向做法效果
Prefill优化切分QK/V,计算与数据读取overlapPrefill持平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
动态阈值稀疏

SpargeAttention
无训练稀疏

SLA
稀疏-线性混合

通过softmax阈值
动态调整块稀疏度

精确且无需训练

可微调,超越
纯Transformer稀疏

BLASST 的核心逻辑:在块级处理中维护全局"运行最大值",每处理一个块,先算块内局部最大值,和全局最大值比较——差值小于阈值就直接跳过这个块(省掉Softmax的指数计算、Value加载和矩阵乘法)。这个过滤过程嵌入在正常计算流水中,零额外开销

加载 Block k

计算 Block k 局部最大值 m_local

m_global - m_local < 阈值?

跳过:不计算Softmax
不加载Value块
不做MatMul

正常计算

更新 m_global

下一个Block

SpargeAttention 的做法不同:先压缩。高相似度的Token块合并成代表性Token,在小尺寸的近似注意力图上判断哪些区域重要,生成稀疏掩码,GPU只计算关键部分。

原始 Q, K

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)

几个值得记住的点:

  1. Online Softmax是分块计算的理论基础,没有增量更新就没有后续的一切。
  2. Chain GEMM的落地关键是Gemm0多次mmac融合,省掉了线程间数据交换,换来18%的性能提升。
  3. 四级Buffer是多级流水的典范,把21KB的LDS用到了极致,加载和计算完全遮掩。
  4. 量化选FP8不选INT8,看的是转换成本——FP32→INT8的转换太贵,FP8有原生Tensor Core撑腰。
  5. 稀疏化三条路线(动态阈值/无训练压缩/可微调混合)解决同一个问题,但切入点完全不同,值得对比着看。

原文中部分涉及内部项目名称和人员信息已做脱敏处理,技术数据和优化思路保持原样。


更多推荐