从零开始写Qwen3(五-其三)FlashAttention提速
概述
前文用triton完成了FlashAttentionV2,成功将显存开销削减到 O ( N ) O(N) O(N),但16位下耗时显著长于官方实现。本文将分步提升性能,最后耗时从官方的237%下降到112%
基准
B=2,N=1024,H=16/8,D=128,使用BM=BN=32, float16
可以逆向推测torch的BM=BN=64,更大的BM和BN会占用更多共享内存和寄存器,并行率会下降,经过测试,至少在现有代码和现有环境上BM=BN=32会更快,但差别不大
torch的耗时为 580us
理论值
理论访存
对于BM=32,一个CTA,需要访问一个BMxD的Q和NxD的KV,以及写入BMxD的O,共 1024 / B M 1024/BM 1024/BM个Block
共读取
32 × 1024 32 ⏟ 线程块数量 × ( 32 × 128 ⏟ Q + 2 × 1024 × 128 ⏟ K V ) × 2 ⏟ f16 = 528 M i B ≈ 566.94 M B \underbrace{32 \times \frac{1024}{32}}_{\text{线程块数量}} \times( \underbrace{32\times 128}_{Q}+\underbrace{2\times 1024\times 128}_{KV})\times \underbrace{2}_{\text{f16}}=528MiB\approx 566.94MB 线程块数量 32×321024×(Q 32×128+KV 2×1024×128)×f16 2=528MiB≈566.94MB
理论写入
32 × 1024 32 × 32 × 128 × 2 = 8 M i B ≈ 8.59 M B 32\times \frac{1024}{32}\times 32\times 128\times 2=8MiB\approx 8.59MB 32×321024×32×128×2=8MiB≈8.59MB
关键指标
- 耗时
- 1380us
- 共享内存43.26K,寄存器128
- 理论占用率16.67,2/12(受限共享内存)
- 带宽
- 126.06G/s
- 推算访存为 126.06 G / s × 1.38 × 10 − 3 s = 173.96 M 126.06G/s\times 1.38\times 10^{-3}s=173.96M 126.06G/s×1.38×10−3s=173.96M
- 总访存
- 173.96M
- 读取164.3M
- 写入9.25M
- L2命中率72.18%
- 读写L2,553.65
- 读取545.26M
- 写入8.39M
- 读写L2,553.65
- 126.06G/s
改进1-使用正确的grid
原始代码中grid先bh再q,这相当于
原来的grid:0为bh,1为q,但是GPU执行顺序是0优先,所以相当于是跑了这个循环
for q in range(0, M, BLOCK_SIZE_M):
for bh in range(BH):
...
这样L2命中率肯定是低下的,改成x为Q,y为BH
数据变化
| 耗时(us) | 全局读取(MB) | L2命中率 | |
|---|---|---|---|
| 绝对 | 1380 | 17.1 | 96.4% |
| 相对 | -0% | -97% | +25.21% |
全局读取极大减少,基本所有的KV都被L2缓存下来,然而耗时根本不减少,说明代码根本不是受限于访存,而是计算
改进2-KV只读取下三角
在因果遮罩中,当n_k ≥ n_q 时,应该直接得到0,我们在代码中的实现是在exp前设置为-inf实现的,但实际上在遍历中
for k in tl.range(0, N, BLOCK_SIZE_N):
对于 k>pid * BLOCK_SIZE_M 的部分可以完全不用计算,因为注意力权重肯定是0,这样可以省去大量读取计算,改成这种
if is_causal:
hi = min(N, (pid + 1) * BLOCK_SIZE_M)
else:
hi = N
for k in tl.range(0, hi, BLOCK_SIZE_N):
k = tl.multiple_of(k, BLOCK_SIZE_N)
理论上可以减少多少呢,对于BM和BN相同,M和N相同的情况,只用计算下三角,循环次数减少:
32 × 32 → 32 × 1 + 32 2 32\times 32\to 32\times \frac{1+32}{2} 32×32→32×21+32
减少48.4375%
实际结果如下
| 耗时(us) | L2读取(MB) | FLOPS(MFLOPS) | |
|---|---|---|---|
| 绝对 | 818 | 285.21 | 144.6 |
| 相对上一步 | -40.72% | -47.69% | -53.04% |
这是相当有效的改进
改进3-两阶段计算
改进2中实际真的设置遮罩的只有对角线部分,非对角线部分可以不做,这就可以把KV循环分成两个部分:没有遮罩的部分和有遮罩的部分
@triton.jit
def fused_attention_intr(
...
):
if STAGE == 1:
lo, hi = 0, min(start_m * BLOCK_SIZE_M, N)
elif STAGE == 2:
lo, hi = min(start_m * BLOCK_SIZE_M, N), min((start_m + 1) * BLOCK_SIZE_M, M)
else:
lo, hi = 0, N
for k in tl.range(lo, hi, BLOCK_SIZE_N):
...
if STAGE == 2:
attn = tl.where(offsets_m[:, None] >= offsets_n[None, :], attn, -float("inf"))
...
return result_o, max_val, dominator, offsets_n
if is_causal:
result_o, max_val, dominator, offsets_n = fused_attention_intr(
...
STAGE=1,
...
)
result_o, max_val, dominator, offsets_n = fused_attention_intr(
...
STAGE=2,
...
)
else:
result_o, max_val, dominator, offsets_n = fused_attention_intr(
...
STAGE=3,
...
)
result_o = (result_o / dominator).to(dtype)
tl.store(
O_ptr + offsets_qm[:, None] * stride_om + offsets_qd[None, :] * stride_od,
result_o.to(dtype),
mask=mask_d & mask_m,
)
优化2中因果遮罩部分占了总指令指令的6.22%,可以估算,这条指令执行了
32 × 32 + 1 2 = 528 32\times \frac{32+1}{2}=528 32×232+1=528
遍,所以每次执行如果分开计算,则只用计算32次,算起来可以减少
6.22 % × ( 1 − 27.5 32 ) ≈ 5 % 6.22\%\times(1-\frac{27.5}{32})\approx5\% 6.22%×(1−3227.5)≈5%
实际结果
| 耗时(us) | |
|---|---|
| 绝对 | 780 |
| 相对上一步 | -4.64% |
改进点4-去掉多余的转换
包括result_o,max_val,dominator全部使用f32,而qkv使用f16,不做多余转换
效果
| 耗时(us) | 寄存器 | 执行指令数 | |
|---|---|---|---|
| 绝对 | 792 | 127 | - |
| 相对上一步 | +1.57% | -6.62% | -7% |
虽然耗时增加了一点点,但f32中间值要保证精度,而且最关键的是减少了寄存器,降到128以下,这意味着只考虑寄存器的情况,一个SM能放的线程块从3升到4,尽管当前还是首先于共享内存限制在2,占用率为 2/12 = 16.67%,但下一步就能看到占用率提升
优化点5-使用warp_specialize
for k in tl.range(lo, hi, BLOCK_SIZE_N, warp_specialize=True):
| 耗时(us) | 共享内存(MB) | 占用率 | |
|---|---|---|---|
| 绝对 | 668 | 18.43 | 4/12 |
| 相对上一步 | -16.23% | -57.14% | +200% |
warp_specialize,可以将线程束划分为负责搬运数据的和负责计算的,可能这样每次只用加载部分共享内存,不用全部加载
优化点6-使用exp2而不是exp
实际上,CUDA中的指数用的是 2 x 2^x 2x,而不是 e x e^x ex,所以tl.exp(x) 实际上是两条指令 x*=math.log2(math.e),exp2(x) ,现在直接在入口传scale时把log2(e)乘上去,然后使用tl.math.exp2
原始两条exp,分别占总指令数的10.02%和5.98%,这个指令可以说热点指令了,减少一个乘法可以减少
16 % × 1 2 = 8 % 16\%\times \frac{1}{2}=8\% 16%×21=8%
的指令,实际减少了11.48%的执行指令,但耗时只减少2.46%
| 耗时(us) | 指令执行数 | |
|---|---|---|
| 绝对 | 654 | - |
| 相对上一步 | -2.46% | -11.48% |
改进点7-把stride当成常量
之前虽然使用了stride,但基本都假设stride_d为1,那舍弃一些灵活性,直接要求stride_d为1,stride_m,n为d常数
减少10%的指令,但耗时只减少0.36%到652us,分支指令降低84%,分支效率提升到100%
差距
torch580us,距离目标还差12%
可能差距的地方
- 执行指令数,多了85%
- 共享内存读取写入多了将近50%
如果使用了 BM=BN=64,耗时升到700us,因为共享内存增加为原来4倍,占用率下降到2/12。
尝试把 scale 提前乘到 Q 上,但几乎没有区别,查看汇编发现triton把这两条指令合并了
attn = attn * scale
attn = attn - new_max_val
成一个FFMA,如果把Scale提前乘到Q,这条指令就退化为 FADD,也没有减少什么东西
更多推荐


所有评论(0)