从零开始写Qwen3目录

概述

前文中写了自注意力的基础版本,成功跑起了推理,但这个实现有个问题,在实际使用中不得不面对:它申请了一块 M × N M\times N M×N 的注意力权重矩阵

当MN变大时,比如prefix阶段计算一个长度为S的提示词,就需要申请 S 2 S^2 S2 大小的注意力权重,这至少会产生三个开销

  • 额外的显存开销
  • 开辟、释放显存的开销
  • 读写全局内存的开销

如果能避免这个中间矩阵申请就能同时降低显存开销和计算开销

简易情况-三矩阵相乘

如果没有这个softmax,三矩阵乘法消掉中间矩阵申请是比较容易的

按照torch的行优先写法,向量都是行向量

把最终结果按行划分

O = [ o 1 o 2 ⋮ o n ] O = \left[\begin{matrix}o_1\\o_2\\\vdots\\ o_n\end{matrix}\right] O= o1o2on

把QK按照行划分,把V按照列分

Q = [ q 1 q 2 ⋮ q m ] , K = [ k 1 k 2 ⋮ k n ] , V = [ v 1 v 2 ⋮ v n ] Q=\left[\begin{matrix}q_1\\q_2 \\\vdots\\q_m\end{matrix}\right] ,K=\left[\begin{matrix}k_1\\k_2\\\vdots\\k_n \end{matrix}\right], V=\left[\begin{matrix}v_1\\v_2\\\vdots\\v_n \end{matrix}\right] Q= q1q2qm ,K= k1k2kn ,V= v1v2vn

暂时忽略dropout和mask,这些都是单点运算,不影响过程,结果表示是

O = [ q i k j ⊤ v j d ] O = \left[\frac{q_i k_j^\top v_j}{\sqrt{d}} \right] O=[d qikjvj]

这样都不需要注意力权重,可以用两次迭代算完

此时只需要共享内存把QKV三个块的分块存下来就行,需要

B M × B K + 2 B N × B K B_M\times B_K + 2 B_N \times B_K BM×BK+2BN×BK

即可,这是一个常数大小的空间,不需要动态长度分配

然而softmax的存在使得它不能简单计算,O变成了

O = [ exp ⁡ ( q i k j ⊤ / d − max ⁡ l ( q i k l ⊤ / d ) ) ∑ k exp ⁡ ( q i k k ⊤ / d − max ⁡ l ( q i k l ⊤ / d ) ) v j ] O = \left[\begin{matrix}\frac{\exp(q_i k^\top_j/\sqrt{d} -\max_l (q_i k_l^\top/\sqrt{d}))}{\sum_k \exp(q_i k_k^\top/\sqrt{d} - \max_l (q_i k_l^\top/\sqrt{d}))} \end{matrix} v_j\right] O=[kexp(qikk/d maxl(qikl/d ))exp(qikj/d maxl(qikl/d ))vj]

这样为了算一个结果,需要用的所有K数据,无法分块计算

在线softmax

考虑一个

v ∈ R D = [ V 1 , … , V n ] , V i ∈ R d v\in \mathbb{R}^D=[V_1,\dots,V_n],V_i \in \mathbb{R}^d vRD=[V1,,Vn],ViRd

要以d分块迭代算它的softmax

首先一个问题就是要除以分母,而分母是 ∑ i exp ⁡ ( v i ) \sum_i \exp(v_i) iexp(vi),这个求和要用到所有数据,在单次迭代中无法获取

另外一个问题就是计算softmax需要用到 max ⁡ ( v ) \max(v) max(v),每个 v i v_i vi 减去这个最大值再计算指数,这也需要用到所有数据

解决办法就是在线更新,也就是每次计算分块d内部的最大值和求和值,得到新的数据,同时需要更新老数据

V i ′ = V i − M i , M i = max ⁡ k < i ( M k ) , M 1 = max ⁡ ( V 1 ) , S i = ∑ k i d exp ⁡ ( v k − M i ) V_i' = V_i - M_i, M_i = \max_{k <i}(M_k), M_1 = \max(V_1), S_i = \sum_k^{id}\exp(v_k - M_i) Vi=ViMi,Mi=k<imax(Mk),M1=max(V1),Si=kidexp(vkMi)

第一段直接计算

O 1 = exp ⁡ ( V 1 ′ ) S 1 − 1 O_1 = \exp(V_1')S_1^{-1} O1=exp(V1)S11

第二段则需要更新得到最新的 M 2 = max ⁡ ( M 1 , max ⁡ ( V 2 ) ) M_2=\max (M_1, \max(V_2)) M2=max(M1,max(V2)),得到 V 2 ′ V'_2 V2

计算

O 2 = exp ⁡ ( V 2 ′ ) S 2 − 1 O_2 = \exp(V_2')S_2^{-1} O2=exp(V2)S21

此时回去更新 V 1 V_1 V1,分子分母都需要更新,分子乘以

exp ⁡ ( M 1 ) exp ⁡ ( M 2 ) \frac{\exp(M_1)}{\exp(M_2)} exp(M2)exp(M1)

这样能把分子修正成

exp ⁡ ( V 1 − M 1 ) exp ⁡ ( M 1 ) exp ⁡ ( M 2 ) = exp ⁡ ( V 1 − M 2 ) \exp(V_1-M_1) \frac{\exp(M_1)}{\exp(M_2)}=\exp(V_1-M_2) exp(V1M1)exp(M2)exp(M1)=exp(V1M2)

分母也需要统一成 S 2 S_2 S2

也就是

O 1 → O 1 exp ⁡ ( M 1 − M 2 ) S 1 S 2 − 1 O_1\to O_1 \exp(M_1-M_2) S_1 S_2^{-1} O1O1exp(M1M2)S1S21

实际上可以把整个O都乘以这个系数(还没有写入的部位是0,乘以系数还是0),所以每次迭代时的公式是这样的

O → O exp ⁡ ( M i − M i + 1 ) S i S i + 1 − 1 + [ 0 ; O i + 1 ; 0 ] O\to O\exp(M_i-M_{i+1})S_{i}S_{i+1}^{-1}+ [\bold{0}; O_{i+1}; \bold{0}] OOexp(MiMi+1)SiSi+11+[0;Oi+1;0]

FlashAttention

有了在线softmax,就可以实现一个不需要额外申请注意力权重矩阵的自注意力,也就是FlashAttention

把QKV按照 B M , B N B_M,B_N BM,BN分块

Q = [ Q 1 Q 2 ⋮ Q m ] , K = [ K 1 K 2 ⋮ K n ] , V = [ V 1 V 2 ⋮ V n ] Q=\left[\begin{matrix}Q_1\\Q_2 \\\vdots\\Q_m\end{matrix}\right] ,K=\left[\begin{matrix}K_1\\K_2\\\vdots\\K_n \end{matrix}\right], V=\left[\begin{matrix}V_1\\V_2\\\vdots\\V_n \end{matrix}\right] Q= Q1Q2Qm ,K= K1K2Kn ,V= V1V2Vn

O的尺寸和Q一致,也是一样划分

然后每个Q_i负责产生一个O_i,遍历所有K和V:

O i = ∑ j softmax ( mask ( Q i K j ⊤ / D ) ) V j , O j ∈ R B M × D O_i = \sum_j \text{softmax}(\text{mask}(Q_i K_j^\top/\sqrt{D}))V_j, O_j \in \mathbb{R}^{B_M\times D} Oi=jsoftmax(mask(QiKj/D ))Vj,OjRBM×D

这里的softmax就可以用在线更新,不过不同的是这里不是更新中间的内容,而是更新的和V相乘后的结果

O i → O i exp ⁡ ( M j − M j + 1 ) S j S j + 1 − 1 + exp ⁡ ( mask ( Q i K j + 1 ⊤ / D ) − M j + 1 ) S j + 1 − 1 V j + 1 O_i\to O_i \exp(M_j - M_{j+1}) S_{j}S_{j+1}^{-1} + \exp(\text{mask}(Q_iK_{j+1}^\top/\sqrt{D}) -M_{j+1})S_{j+1}^{-1}V_{j+1} OiOiexp(MjMj+1)SjSj+11+exp(mask(QiKj+1/D )Mj+1)Sj+11Vj+1

用代码实现就是

@triton.jit
def fused_attention(
    Q,
    K,
    V,
    output,
    M,
    N,
    d,
    scale,
    stride_qh,
    stride_qm,
    stride_qd,
    stride_kh,
    stride_kn,
    stride_kd,
    stride_vh,
    stride_vn,
    stride_vd,
    stride_oh,
    stride_om,
    stride_od,
    is_causal: tl.constexpr,
    groups: tl.constexpr,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr,
    USE_FP32_ACCUM: tl.constexpr,
    dtype: tl.constexpr,
):
    pid_qh = tl.program_id(0)  # BxH ,放到一个编号中
    pid_kh = (pid_qh) // groups  # KV 的编号向下整除 组数,不用分离 B和H,自然和Q在同一个B中
    pid = tl.program_id(1)
    result_o = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_K), dtype=tl.float32)
    offsets_qm = pid * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
    offsets_qd = tl.arange(0, BLOCK_SIZE_K)

    Q_ptr = Q + pid_qh * stride_qh
    K_ptr = K + pid_kh * stride_kh
    V_ptr = V + pid_kh * stride_vh
    O_ptr = output + pid_qh * stride_oh

    mask_m = offsets_qm[:, None] < M
    mask_d = offsets_qd < d

    data_q = tl.load(
        Q_ptr + offsets_qm[:, None] * stride_qm + offsets_qd[None, :] * stride_qd,
        mask=mask_m & mask_d,
        other=0.0,
    ).to(dtype)

    max_val = tl.zeros((BLOCK_SIZE_M, 1), dtype=dtype) - float("inf")
    dominator = tl.zeros((BLOCK_SIZE_M, 1), dtype=dtype)
    for k in tl.range(0, N, BLOCK_SIZE_N):
        offsets_kk = k + tl.arange(0, BLOCK_SIZE_N)
        kv_mask = mask_d & (offsets_kk[:, None] < N)
        data_k = tl.load(
            K_ptr + offsets_kk[:, None] * stride_kn + offsets_qd[None, :] * stride_kd,
            mask=kv_mask,
            other=0.0,
        ).to(dtype)
        data_v = tl.load(
            V_ptr + offsets_kk[:, None] * stride_vn + offsets_qd[None, :] * stride_vd,
            mask=kv_mask,
            other=0.0,
        ).to(dtype) 
        if USE_FP32_ACCUM:
            attn = tl.dot(data_q, data_k.T, input_precision="ieee") * scale
        else:
            attn = tl.dot(data_q, data_k.T) * scale
        attn = tl.where(mask_m & (offsets_kk[None, :] < N), attn, -float("inf"))
        if is_causal:
            attn = tl.where(offsets_qm[:, None] >= offsets_kk[None, :], attn, -float("inf"))
        tmp_max = tl.max(attn, axis=-1, keep_dims=True)
        new_max_val = tl.maximum(max_val, tmp_max)
        attn = attn - new_max_val
        exp_attn = tl.exp(attn)

        scale_factor = tl.exp(max_val - new_max_val)
        dominator = dominator * scale_factor + tl.sum(exp_attn, axis=-1, keep_dims=True)
        max_val = new_max_val
        if USE_FP32_ACCUM:
            result_o = result_o * scale_factor + tl.dot(exp_attn, data_v, input_precision="ieee")
        else:
            result_o = result_o * scale_factor + tl.dot(exp_attn, data_v)
    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,
    )

这里通过常量 USE_FP32_ACCUM 来判断是否要保留点乘的精度,在测试用例里面要用,生成就不用了。用is_causal 来控制是否需要增加因果遮罩

这里矩阵乘法简单全部.to(dtype),后面可以研究如何更好地处理数据类型,比如量化情况

变体

最开始FlashAttn的作者认为每个Q_i都要重复读取所有K_j和V_j,有重复,于是只有BH并行,KV作为外循环,Q作为内循环,这是FlashAttnV1

然后后来发现重复读取其实也不是那么遭,毕竟有缓存在,反而是Seqlen较长时不能充分利用并行能力,于是对Q按照SeqLen进行并行,得到FlashAttnV2,也就是上面的写法,Q并行,KV作为内循环

实验对比

3060-12G上执行,B=2,float32

base使用torch.nn.functional.scaled_dot_product_attention,my_op是分三步的算子,my_op_flash是flashattn
在这里插入图片描述

scaled_dot_product_attention是Torch官方的多头注意力实现。my_op申请了attn,和base一样,它的最大显存都是 O ( N 2 ) O(N^2) O(N2),但相比torch的python实现,triton的显存开销少了大概一半(已经开启no_grad,所以没有梯度问题),而my_op_flash使用了FlashAttn,显存是线性增长的

2048长度时,FlashAttn的内存开销是32M,刚好就是输出的大小( 2 ( B ) × 2048 ( S ) × 2024 ( D ) × 4 ( B / float ) 2(B)\times 2048(\text{S})\times 2024(\text{D}) \times 4(B/\text{float}) 2(B)×2048(S)×2024(D)×4(B/float)),FlashAttn没有申请任何额外空间

在这里插入图片描述

FlashAttn不仅显存开销小,速度也更块。在这个测试中seq_len不能再增长了,因为12G显存已经不够 O ( N 2 ) O(N^2) O(N2) 了,triton的测试占用显存有点过大

在这里插入图片描述

单独FlashAttn可以更长,还远没有到极限

补充实验

FP16的scaled_dot_product_attention用的应该就是FlashAttn
在这里插入图片描述

--------------------------------------------------------------------------------
   Seq Len |    base (MB) |   my_op (MB) |   my_op_flash (MB)
--------------------------------------------------------------------------------
        64 |         0.51 |         0.75 |               0.50
       128 |         1.02 |         2.00 |               1.00
       192 |         1.52 |         3.75 |               1.50

在这里插入图片描述
并且Torch的实现比当前实现快很多,下一步就是要定位原因并提速接近Torch

Logo

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

更多推荐