从零开始写Qwen3(五-其二)使用Triton实现FlashAttention
概述
前文中写了自注意力的基础版本,成功跑起了推理,但这个实现有个问题,在实际使用中不得不面对:它申请了一块 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= o1o2⋮on
把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= q1q2⋮qm ,K= k1k2⋮kn ,V= v1v2⋮vn
暂时忽略dropout和mask,这些都是单点运算,不影响过程,结果表示是
O = [ q i k j ⊤ v j d ] O = \left[\frac{q_i k_j^\top v_j}{\sqrt{d}} \right] O=[dqikj⊤vj]
这样都不需要注意力权重,可以用两次迭代算完
此时只需要共享内存把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 v∈RD=[V1,…,Vn],Vi∈Rd
要以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′=Vi−Mi,Mi=k<imax(Mk),M1=max(V1),Si=k∑idexp(vk−Mi)
第一段直接计算
O 1 = exp ( V 1 ′ ) S 1 − 1 O_1 = \exp(V_1')S_1^{-1} O1=exp(V1′)S1−1
第二段则需要更新得到最新的 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′)S2−1
此时回去更新 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(V1−M1)exp(M2)exp(M1)=exp(V1−M2)
分母也需要统一成 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} O1→O1exp(M1−M2)S1S2−1
实际上可以把整个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}] O→Oexp(Mi−Mi+1)SiSi+1−1+[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= Q1Q2⋮Qm ,K= K1K2⋮Kn ,V= V1V2⋮Vn
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=j∑softmax(mask(QiKj⊤/D))Vj,Oj∈RBM×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} Oi→Oiexp(Mj−Mj+1)SjSj+1−1+exp(mask(QiKj+1⊤/D)−Mj+1)Sj+1−1Vj+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
更多推荐


所有评论(0)