从零开始写Qwen3目录

概述

前文用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=528MiB566.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=8MiB8.59MB

关键指标

  1. 耗时
    1. 1380us
  2. 共享内存43.26K,寄存器128
    1. 理论占用率16.67,2/12(受限共享内存)
  3. 带宽
    1. 126.06G/s
      1. 推算访存为 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×103s=173.96M
    2. 总访存
      1. 173.96M
      2. 读取164.3M
      3. 写入9.25M
    3. L2命中率72.18%
      1. 读写L2,553.65
        1. 读取545.26M
        2. 写入8.39M

改进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×3232×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%×(13227.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,也没有减少什么东西

Logo

欢迎加入DeepSeek 技术社区。在这里,你可以找到志同道合的朋友,共同探索AI技术的奥秘。

更多推荐