Torch朴素实现
定义

MLA的主要思想就是,把KV压缩到隐空间,计算时再解压,这样存在显存的KV Cache只用存隐空间向量,节约显存
Attention Score=Q⋅KT=Q⋅(cKVWUK)T=(Q⋅(WUK)T)⋅(cKV)T\text{Attention Score} = Q \cdot K^T = Q \cdot (c^{KV} W^{UK})^T = (Q \cdot (W^{UK})^T) \cdot (c^{KV})^TAttention Score=Q⋅KT=Q⋅(cKVWUK)T=(Q⋅(WUK)T)⋅(cKV)T
为了解决位置编码的问题,,将QK向量分成两部分,一部分是不含位置编码的,另一部分单独拿出来做位置编码,K的位置编码部分不被压缩到隐空间,将位置编码部分和原始向量拼接,得到完整QK向量。
下面以一个具体torch实现为例,讲解MLA的朴素实现流程,搞清楚定义了,后面再看Tilelang的优化实现
输入reshape
q = rearrange(q, "b (h g) d -> b g h d", g=num_head_groups)将head维度拆成(group,head),类似GQA的思路,但是由于本题的数据实际上都是KV只有一个头,也就是g=1,所以没什么影响。这里不用torch.reshape,用这个einops.rearrange,可以传入一个字符串,给每个维度起名字,然后用字符的排列表示想实现的reshape效果,更直观。
q_pe = rearrange(q_pe, "b (h g) d -> b g h d", g=num_head_groups)mla将位置编码和和内容分开,所以位置编码是单独一段
kv = rearrange(kv, "b n h d -> b h n d")kv需要把最后两个维度变成(seqlen,dim)
k_pe = rearrange(k_pe, "b n h d -> b h n d")k也有位置编码
py
# 先显式展开"同一个 kv head 对应多个 query head"的 group 维度。
q = rearrange(q, "b (h g) d -> b g h d", g=num_head_groups)
q_pe = rearrange(q_pe, "b (h g) d -> b g h d", g=num_head_groups)
# 再把 kv 调整成 [batch, kv_head, seq, dim] 的形式。
kv = rearrange(kv, "b n h d -> b h n d")
k_pe = rearrange(k_pe, "b n h d -> b h n d")
拼接内容和位置编码
query = torch.concat([q, q_pe], dim=-1)根据原始定义朴素实现,把位置编码和原始向量在最后一个维度(dim)直接拼接。
py
# 参考实现直接显式拼接内容维和位置维。
# 这与 kernel 里"分两次 gemm 再相加"的数学结果完全一致。
query = torch.concat([q, q_pe], dim=-1)
key = torch.concat([kv, k_pe], dim=-1)
计算注意力矩阵
用拼接后的qk计算注意力得分,也可以用这个einops库,不用指定转置,会根据我们传入的shape自动推导。这里注意到d维度消失了,那么就能推导出是d维度做点积了。
py
# 得到每个 query head 对整段 KV 序列的注意力分数。
scores = einsum(query, key, "b g h d, b h s d -> b g h s")
attention = F.softmax(scores / scale, dim=-1)
乘上V得到注意力输出
out = einsum(attention, kv, "b g h s, b h s d -> b g h d")注意力矩阵再和V做加权求和,得到注意力输出。这里k,v实际用的都是KV张量,这是MLA的优化,KV压缩到隐空间,共享一个cache张量,而不是K,V分别一个张量,KV的区分是通过解压矩阵来实现的。
out = rearrange(out, "b g h d -> b (h g) d")最后reshape成和输入一样的格式
py
# value 端仍然只使用 kv 内容部分。
# 位置编码不进入输出聚合,这一点与 MLA 设计保持一致。
out = einsum(attention, kv, "b g h s, b h s d -> b g h d")
out = rearrange(out, "b g h d -> b (h g) d")
K解压
根据MLA的定义,可能会疑惑KV cache不是需要分别解压才能得到K,V吗,解压去哪了?
一般来讲需要解压K=ctKVWDKK = c_t^{KV} W^{DK}K=ctKVWDK但是带入注意力公式可得Scores=QKT=Q(ctKVWDK)T=Q(WDK)T(ctKV)T\text{Scores} = Q K^T = Q (c_t^{KV} W^{DK})^T = Q (W^{DK})^T (c_t^{KV})^TScores=QKT=Q(ctKVWDK)T=Q(WDK)T(ctKV)T利用矩阵乘法的结合律,我们可以把 (WDK)T(W^{DK})^T(WDK)T 提前和 QQQ 乘起来:Scores=(Q(WDK)T)(ctKV)T\text{Scores} = \left( Q (W^{DK})^T \right) (c_t^{KV})^TScores=(Q(WDK)T)(ctKV)T也就是我们可以把解压矩阵提前和Q相乘,这里传入的Q矩阵可能就已经包含解压矩阵了。
V解压
为什么V也不需要解压?标准 MLA 中,从隐变量得到 VVV 需要乘以解压矩阵 WDVW^{DV}WDV:V=ctKVWDVV = c_t^{KV} W^{DV}V=ctKVWDV
最终注意力机制的输出(不考虑 RoPE 部分)为:O=Attention×V=Attention×(ctKVWDV)\text{O} = \text{Attention} \times V = \text{Attention} \times (c_t^{KV} W^{DV})O=Attention×V=Attention×(ctKVWDV)
再次利用结合律,把 WDVW^{DV}WDV 提出来移到最后:O=(Attention×ctKV)WDV\text{O} = \left( \text{Attention} \times c_t^{KV} \right) W^{DV}O=(Attention×ctKV)WDV
也就是可以在这个函数返回后,再乘上V解压矩阵
完整代码
py
# 这是纯 PyTorch 的参考实现。
# 它主要用于 correctness check,不追求任何 kernel 级优化。
def ref_program(q, q_pe, kv, k_pe):
"""
输入张量形状如下。
q 的形状是 [batch, heads, dim]。
q_pe 的形状是 [batch, heads, pe_dim]。
kv 的形状是 [batch, seqlen_kv, kv_head_num, dim]。
k_pe 的形状是 [batch, seqlen_kv, kv_head_num, pe_dim]。
输出 out 的形状是 [batch, heads, dim]。
"""
dim = q.shape[-1]
pe_dim = q_pe.shape[-1]
num_head_groups = q.shape[1] // kv.shape[2]
scale = (dim + pe_dim) ** 0.5
# 先显式展开"同一个 kv head 对应多个 query head"的 group 维度。
q = rearrange(q, "b (h g) d -> b g h d", g=num_head_groups)
q_pe = rearrange(q_pe, "b (h g) d -> b g h d", g=num_head_groups)
# 再把 kv 调整成 [batch, kv_head, seq, dim] 的形式。
kv = rearrange(kv, "b n h d -> b h n d")
k_pe = rearrange(k_pe, "b n h d -> b h n d")
# 参考实现直接显式拼接内容维和位置维。
# 这与 kernel 里"分两次 gemm 再相加"的数学结果完全一致。
query = torch.concat([q, q_pe], dim=-1)
key = torch.concat([kv, k_pe], dim=-1)
# 得到每个 query head 对整段 KV 序列的注意力分数。
scores = einsum(query, key, "b g h d, b h s d -> b g h s")
attention = F.softmax(scores / scale, dim=-1)
# value 端仍然只使用 kv 内容部分。
# 位置编码不进入输出聚合,这一点与 MLA 设计保持一致。
out = einsum(attention, kv, "b g h s, b h s d -> b g h d")
out = rearrange(out, "b g h d -> b (h g) d")
return out
Tilelang no-split实现
Tilelang也有简单实现和复杂实现的区别,先看简单实现。no-split指的是,kv cache也就是seqlen可能很长,一个优化思路是对seqlen分块,分给不同线程块,但是这会引入更多的分块维度,太复杂。先看无seqlen分块的实现
分块
首先是最重要的分块模式,对q的head维度分块,每块负责一些头。再对batch维度分块,每块负责一个q batch和一个kv batch,是一个二维分块。
py
with T.Kernel(heads // min(block_H, kv_group_num), batch, threads=128) as (hid, bid)
申请块内存
申请分块临时内存,block_H表示head维度分块大小,block_N表示seqlen维度分块大小,注意前面说的不对seqlen分块,是指线程块维度不分,一个线程块负责一个序列的全部kv,但是实际计算时还是划窗或者说分块计算的,只不过都由一个线程块串行计算。
S_shared 保存注意力得分的中间结果,所以shape是(block_H, block_N),保存最终的注意力输出,乘上V之后shape回到(block_H, dim)。Q_shared, Q_pe_shared ,KV_shared ,K_pe_shared 分别保存取出来计算的Q,K tile。
acc_s,acc_o是寄存器张量,用于gemm计算时累加结果,因为矩阵乘法累加会多次访问目的地址,所以选最快的寄存器
scores_max 这些都是online softmax需要的中间变量,sum,max每行只需要一个变量,所以shape都是(block_H)。
py
Q_shared = T.alloc_shared([block_H, dim], dtype)
S_shared = T.alloc_shared([block_H, block_N], dtype)
Q_pe_shared = T.alloc_shared([block_H, pe_dim], dtype)
KV_shared = T.alloc_shared([block_N, dim], dtype)
K_pe_shared = T.alloc_shared([block_N, pe_dim], dtype)
O_shared = T.alloc_shared([block_H, dim], dtype)
# acc_s 保存当前 score tile。
acc_s = T.alloc_fragment([block_H, block_N], accum_dtype)
# acc_o 保存扫描完整段 KV 后的输出累加值。
acc_o = T.alloc_fragment([block_H, dim], accum_dtype)
# 下面这些寄存器共同维护 online softmax 的状态。
scores_max = T.alloc_fragment([block_H], accum_dtype)
scores_max_prev = T.alloc_fragment([block_H], accum_dtype)
scores_scale = T.alloc_fragment([block_H], accum_dtype)
scores_sum = T.alloc_fragment([block_H], accum_dtype)
logsum = T.alloc_fragment([block_H], accum_dtype)
读入Q+初始化张量
这个算子是decode阶段,所以q的seqlen为1也就是只有一个q,虽然还有个head维度,但head维度我们分块了,现在q只有一个确定的tile,后面和kv计算的都是这个tile,直接提前读入。
需要累加的张量,置零。需要维护最大值的,置为负无穷。
py
# 先把 query 内容和位置编码读入 shared memory。
T.copy(Q[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, :], Q_shared)
T.copy(Q_pe[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, :], Q_pe_shared)
T.fill(acc_o, 0)
T.fill(logsum, 0)
T.fill(scores_max, -T.infinity(accum_dtype))
计算注意力得分+online softmax第一轮循环
T.copy(KV[bid, k * block_N : (k + 1) * block_N, cur_kv_head, :], KV_shared)先读入KV和k_pe位置编码
T.gemm(Q_shared, KV_shared, acc_s, transpose_B=True, policy=T.GemmWarpPolicy.FullCol, clear_accum=True)分别和q,q_pe点乘计算注意力得分。
两次矩阵乘法都累加到acc_s,中间不清零,相当于分别计算完加起来。这和先拼接k k_pe,q q_pe再乘是等价的,因为拼接的是dim维度,dim维度做点乘最后变成一个标量了,那先拼接再点乘和分别点乘,再标量累加是等价的。
后面这段就是对注意力得分做online softmax了,第一轮循环先算出来sum和max
T.copy(scores_max, scores_max_prev)先把前缀最值保存,T.fill(scores_max, -T.infinity(accum_dtype))接着置为负无穷,用来计算当前块内的每行最值,计算最值直接T.reduce_max(acc_s, scores_max, dim=1, clear=False)调用reduce接口,接着对于每一行,和前缀最值对比,更新目前为止的全局最值scores_max[i] = T.max(scores_max[i], scores_max_prev[i])
最值更新后是online softmax的核心,用新的最值缩放之前的sum,scores_scale[i] = T.exp2(scores_max_prev[i] * scale - scores_max[i] * scale),这里先计算缩放比例,保存
acc_s[i, j] = T.exp2(acc_s[i, j] * scale - scores_max[i] * scale)对当前块的注意力得分做dk\sqrt{d_k}dk 缩放和减去max
T.reduce_sum(acc_s, scores_sum, dim=1)累加这一块每一行的指数和,logsum[i] = logsum[i] * scores_scale[i] + scores_sum[i]用当前块指数和和缩放因子,更新指数和
acc_o[i, j] *= scores_scale[i]acc_o是eqkTve^{qk^T}veqkTv,相比完整的注意力还差一个softmax缩放,所以也需要累加过程中根据max缩放
T.gemm(S_shared, KV_shared, acc_o, policy=T.GemmWarpPolicy.FullCol)缩放后,当前块的结果也累加进来
py
# 对完整 KV 序列做单次线性扫描。
# 每次取一个 block_N 大小的 tile 进行打分、softmax 和 value 聚合。
loop_range = T.ceildiv(seqlen_kv, block_N)
for k in T.Pipelined(loop_range, num_stages=0):
T.copy(KV[bid, k * block_N : (k + 1) * block_N, cur_kv_head, :], KV_shared)
T.copy(K_pe[bid, k * block_N : (k + 1) * block_N, cur_kv_head, :], K_pe_shared)
# 当前 tile 的分数由内容项和位置项两部分组成。
T.gemm(Q_shared, KV_shared, acc_s, transpose_B=True, policy=T.GemmWarpPolicy.FullCol, clear_accum=True)
T.gemm(Q_pe_shared, K_pe_shared, acc_s, transpose_B=True, policy=T.GemmWarpPolicy.FullCol)
# 下面这段 online softmax 更新逻辑与 split 版本相同。
T.copy(scores_max, scores_max_prev)
T.fill(scores_max, -T.infinity(accum_dtype))
T.reduce_max(acc_s, scores_max, dim=1, clear=False)
for i in T.Parallel(block_H):
scores_max[i] = T.max(scores_max[i], scores_max_prev[i])
for i in T.Parallel(block_H):
scores_scale[i] = T.exp2(scores_max_prev[i] * scale - scores_max[i] * scale)
for i, j in T.Parallel(block_H, block_N):
acc_s[i, j] = T.exp2(acc_s[i, j] * scale - scores_max[i] * scale)
T.reduce_sum(acc_s, scores_sum, dim=1)
# 这里直接把 softmax 后的 score 放进 shared memory。
# no-split 路径不需要额外的 cast fragment 中转。
T.copy(acc_s, S_shared)
for i in T.Parallel(block_H):
logsum[i] = logsum[i] * scores_scale[i] + scores_sum[i]
for i, j in T.Parallel(block_H, dim):
acc_o[i, j] *= scores_scale[i]
# 用当前 tile 的 softmax 权重与 KV value 相乘,并累加到输出。
T.gemm(S_shared, KV_shared, acc_o, policy=T.GemmWarpPolicy.FullCol)
online softmax第二次循环+拷贝到输出
acc_o[i, j] /= logsum[i]用sum缩放结果,这里没有去缩放注意力矩阵,而是直接去缩放乘上V的注意力输出,因为先缩放注意力矩阵,后面也还是把注意力矩阵当成权重,乘上V得到输出,根据分配率可以直接缩放注意力输出,计算量还小
T.copy(O_shared, Output[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, :])结果拷贝到全局输出张量
py
# 整段 KV 扫描完后,除以最终分母即可得到标准 softmax 输出。
for i, j in T.Parallel(block_H, dim):
acc_o[i, j] /= logsum[i]
T.copy(acc_o, O_shared)
T.copy(O_shared, Output[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, :])
完整代码
py
# no-split 版本不切分 KV。
# 它直接对完整序列做单遍扫描,是更标准的 FlashAttention 形式。
@T.prim_func
def main_no_split(
Q: T.Tensor([batch, heads, dim], dtype),
Q_pe: T.Tensor([batch, heads, pe_dim], dtype),
KV: T.Tensor([batch, seqlen_kv, kv_head_num, dim], dtype),
K_pe: T.Tensor([batch, seqlen_kv, kv_head_num, pe_dim], dtype),
Output: T.Tensor([batch, heads, dim], dtype),
):
# no-split 版本的 grid 是 (head_block, batch)。
with T.Kernel(heads // min(block_H, kv_group_num), batch, threads=128) as (hid, bid):
# 这一组缓存和 split 版本基本一致。
# 只是这里不需要额外保存跨 split 的中间结果。
Q_shared = T.alloc_shared([block_H, dim], dtype)
S_shared = T.alloc_shared([block_H, block_N], dtype)
Q_pe_shared = T.alloc_shared([block_H, pe_dim], dtype)
KV_shared = T.alloc_shared([block_N, dim], dtype)
K_pe_shared = T.alloc_shared([block_N, pe_dim], dtype)
O_shared = T.alloc_shared([block_H, dim], dtype)
# acc_s 保存当前 score tile。
acc_s = T.alloc_fragment([block_H, block_N], accum_dtype)
# acc_o 保存扫描完整段 KV 后的输出累加值。
acc_o = T.alloc_fragment([block_H, dim], accum_dtype)
# 下面这些寄存器共同维护 online softmax 的状态。
scores_max = T.alloc_fragment([block_H], accum_dtype)
scores_max_prev = T.alloc_fragment([block_H], accum_dtype)
scores_scale = T.alloc_fragment([block_H], accum_dtype)
scores_sum = T.alloc_fragment([block_H], accum_dtype)
logsum = T.alloc_fragment([block_H], accum_dtype)
# 计算当前 head block 对应的 kv head。
cur_kv_head = hid // (kv_group_num // block_H)
# 先把 query 内容和位置编码读入 shared memory。
T.copy(Q[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, :], Q_shared)
T.copy(Q_pe[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, :], Q_pe_shared)
T.fill(acc_o, 0)
T.fill(logsum, 0)
T.fill(scores_max, -T.infinity(accum_dtype))
# 对完整 KV 序列做单次线性扫描。
# 每次取一个 block_N 大小的 tile 进行打分、softmax 和 value 聚合。
loop_range = T.ceildiv(seqlen_kv, block_N)
for k in T.Pipelined(loop_range, num_stages=0):
T.copy(KV[bid, k * block_N : (k + 1) * block_N, cur_kv_head, :], KV_shared)
T.copy(K_pe[bid, k * block_N : (k + 1) * block_N, cur_kv_head, :], K_pe_shared)
# 当前 tile 的分数由内容项和位置项两部分组成。
T.gemm(Q_shared, KV_shared, acc_s, transpose_B=True, policy=T.GemmWarpPolicy.FullCol, clear_accum=True)
T.gemm(Q_pe_shared, K_pe_shared, acc_s, transpose_B=True, policy=T.GemmWarpPolicy.FullCol)
# 下面这段 online softmax 更新逻辑与 split 版本相同。
T.copy(scores_max, scores_max_prev)
T.fill(scores_max, -T.infinity(accum_dtype))
T.reduce_max(acc_s, scores_max, dim=1, clear=False)
for i in T.Parallel(block_H):
scores_max[i] = T.max(scores_max[i], scores_max_prev[i])
for i in T.Parallel(block_H):
scores_scale[i] = T.exp2(scores_max_prev[i] * scale - scores_max[i] * scale)
for i, j in T.Parallel(block_H, block_N):
acc_s[i, j] = T.exp2(acc_s[i, j] * scale - scores_max[i] * scale)
T.reduce_sum(acc_s, scores_sum, dim=1)
# 这里直接把 softmax 后的 score 放进 shared memory。
# no-split 路径不需要额外的 cast fragment 中转。
T.copy(acc_s, S_shared)
for i in T.Parallel(block_H):
logsum[i] = logsum[i] * scores_scale[i] + scores_sum[i]
for i, j in T.Parallel(block_H, dim):
acc_o[i, j] *= scores_scale[i]
# 用当前 tile 的 softmax 权重与 KV value 相乘,并累加到输出。
T.gemm(S_shared, KV_shared, acc_o, policy=T.GemmWarpPolicy.FullCol)
# 整段 KV 扫描完后,除以最终分母即可得到标准 softmax 输出。
for i, j in T.Parallel(block_H, dim):
acc_o[i, j] /= logsum[i]
T.copy(acc_o, O_shared)
T.copy(O_shared, Output[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, :])
Tilelang split实现
申请全局内存
增加了对seqlen这个维度的分块,块数num_split,glse保存每个batch块,head块,split块的sum结果
类似地,对于每个split块保存局部的注意力输出Output_partial
py
# glse 保存每个 split 的 log-sum-exp。
# 第二阶段合并不同 split 时,需要依赖这个量重新构造全局 softmax 分母。
glse = T.alloc_global([batch, heads, num_split], dtype)
# Output_partial 保存每个 split 内部已经归一化后的局部输出。
# 第二阶段会用全局 softmax 权重把这些局部输出重新加权求和。
Output_partial = T.alloc_global([batch, heads, num_split, dim], dtype)
分块
这里和no-split比,仍然有batch和head分块,增加了一个seqlen维度的分块num_split,最终是一个三维分块,在seqlen很长的时候,这样可以增加并行度,提高性能
py
# 第一阶段的 grid 是 (batch, head_block, split_id)。
# 每个 CTA 负责一个 batch、一个 query head block、一个 KV split。
with T.Kernel(batch, heads // min(block_H, kv_group_num), num_split, threads=256) as (bid, hid, bz):
拷贝Q和初始化
仍然和no split差不多,Q的seqlen=1,所以一个q和多个kv计算注意力得分,可以在kv循环开始前复制一次Q就行。
用于累加,求sum,max的张量需要初始化
增加的一个是T.use_swizzle(10),这是规定共享内存的映射方式,减少bank conflict,展开讲很复杂,不是这里的重点。只要能理解这是一个内存映射,优化访存就行。感兴趣的可以去搜bank conflict和swizzle
py
# cur_kv_head 表示当前 head block 对应的 kv head。
# 因为这里限定 kv_head_num == 1,所以本质上总会映射到同一个 kv head。
cur_kv_head = hid // (kv_group_num // block_H)
# swizzle 用于优化线程块访问模式,尽量改善 shared memory 和全局内存行为。
T.use_swizzle(10)
# query 会反复与多个 KV tile 相乘。
# 先把对应的 Q 和 Q_pe 搬进 shared memory,可以减少重复读取。
T.copy(Q[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, :], Q_shared)
T.copy(Q_pe[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, :], Q_pe_shared)
# 初始时还没有累积任何输出。
T.fill(acc_o, 0)
# 初始分母设为 0。
T.fill(logsum, 0)
# 初始行最大值设为负无穷,方便后面逐 tile 更新。
T.fill(scores_max, -T.infinity(accum_dtype))
申请块内存
和no-split差不多,基本就是保存Q,K的内存,矩阵乘法累加的内存,做Online softmax的max,sum内存
py
# Q_shared 缓存当前 CTA 负责的 query 内容部分。
Q_shared = T.alloc_shared([block_H, dim], dtype)
# S_shared 缓存 softmax 后的 score tile。
# 它会作为后续 value 侧 gemm 的输入。
S_shared = T.alloc_shared([block_H, block_N], dtype)
# Q_pe_shared 缓存 query 的位置编码部分。
Q_pe_shared = T.alloc_shared([block_H, pe_dim], dtype)
# KV_shared 缓存当前 block 的内容部分。
KV_shared = T.alloc_shared([block_N, dim], dtype)
# K_pe_shared 缓存当前 block 的位置编码部分。
K_pe_shared = T.alloc_shared([block_N, pe_dim], dtype)
# O_shared 用于把寄存器中的输出暂存后再写回全局内存。
O_shared = T.alloc_shared([block_H, dim], dtype)
# acc_s 是 score tile 的累加寄存器。
# 它承接"内容项 gemm"和"位置项 gemm"两次累加。
acc_s = T.alloc_fragment([block_H, block_N], accum_dtype)
# acc_s_cast 用于把 softmax 权重降回 fp16。
# 这样下一步做 value 侧 gemm 时更容易命中高吞吐路径。
acc_s_cast = T.alloc_fragment([block_H, block_N], dtype)
# acc_o 保存输出向量的累加结果。
# 它会随着每个 KV tile 的处理不断更新。
acc_o = T.alloc_fragment([block_H, dim], accum_dtype)
# scores_max 保存每一行当前见过的最大 score。
scores_max = T.alloc_fragment([block_H], accum_dtype)
# scores_max_prev 保存上一轮迭代的最大 score。
# 这样可以做 online softmax 的重缩放。
scores_max_prev = T.alloc_fragment([block_H], accum_dtype)
# scores_scale 是 old_max 切换到 new_max 后的重缩放因子。
scores_scale = T.alloc_fragment([block_H], accum_dtype)
# scores_sum 保存当前 tile 的 softmax 分子和。
scores_sum = T.alloc_fragment([block_H], accum_dtype)
# logsum 保存到当前 tile 为止累计出来的 softmax 分母。
logsum = T.alloc_fragment([block_H], accum_dtype)
注意力计算+online softmax第一次循环
和no-split类似,区别在于,循环范围是,先考虑split分块,再考虑在当前块上block_N步长划窗,也就是循环边界这里T.ceildiv((seqlen_kv // num_split), block_N)
中间的计算过程,区别在于拷贝的数据的seq下标,需要特殊计算,和之前no split不同,之前是一个block就访问全部的seqlen,这里是只访问一个区间,所以要考虑这个区间的起始地址,在这个地址基础上循环偏移kv_start = (seqlen_kv // num_split) * bz + k * block_N
剩下的都是在块内存也就是tile上操作的,数据块计算逻辑完全不变,也不用考虑全局偏移量,所以没有变化
py
# 当前 split 只覆盖整段 KV 中的一部分长度。
# 这里继续按 block_N 做更细粒度分块,并以流水方式迭代。
loop_range = T.ceildiv((seqlen_kv // num_split), block_N)
for k in T.Pipelined(loop_range, num_stages=2):
kv_start = (seqlen_kv // num_split) * bz + k * block_N
kv_end = (seqlen_kv // num_split) * bz + (k + 1) * block_N
# 把当前 KV tile 的内容和位置编码一并读入 shared memory。
T.copy(KV[bid, kv_start:kv_end, cur_kv_head, :], KV_shared)
T.copy(K_pe[bid, kv_start:kv_end, cur_kv_head, :], K_pe_shared)
# 先清空 score 累加器。
T.clear(acc_s)
# 第一项是内容部分的点积。
T.gemm(Q_shared, KV_shared, acc_s, transpose_B=True, policy=T.GemmWarpPolicy.FullCol)
# 第二项是位置编码部分的点积。
# 两次 gemm 的结果相加后,数学上等价于先 concat 再做一次大矩阵乘。
T.gemm(Q_pe_shared, K_pe_shared, acc_s, transpose_B=True, policy=T.GemmWarpPolicy.FullCol)
# 下面进入 online softmax。
# 核心思想是边扫描 KV,边更新最大值和分母,而不是保存完整 score 矩阵。
T.copy(scores_max, scores_max_prev)
T.fill(scores_max, -T.infinity(accum_dtype))
T.reduce_max(acc_s, scores_max, dim=1, clear=False)
# 当前 tile 的最大值要与历史最大值合并,形成新的全局行最大值。
for i in T.Parallel(block_H):
scores_max[i] = T.max(scores_max[i], scores_max_prev[i])
# 历史分母和历史输出都需要按 max 的变化做重缩放。
for i in T.Parallel(block_H):
scores_scale[i] = T.exp2(scores_max_prev[i] * scale - scores_max[i] * scale)
# 把 score 转成 exp(score - max) 形式,得到当前 tile 的 softmax 分子。
for i, j in T.Parallel(block_H, block_N):
acc_s[i, j] = T.exp2(acc_s[i, j] * scale - scores_max[i] * scale)
# 累加每一行在当前 tile 上的分子和。
T.reduce_sum(acc_s, scores_sum, dim=1)
# 先写入 shared memory,再转成 fp16 fragment。
# 这是为了兼顾 softmax 阶段的精度和 gemm 阶段的吞吐。
T.copy(acc_s, S_shared)
T.copy(S_shared, acc_s_cast)
# 用重缩放后的旧分母加上当前 tile 的新分子和,得到新的分母。
for i in T.Parallel(block_H):
logsum[i] = logsum[i] * scores_scale[i] + scores_sum[i]
# 旧输出也必须同步按同样比例缩放。
# 这是 online softmax 保持数值正确的关键步骤。
for i, j in T.Parallel(block_H, dim):
acc_o[i, j] *= scores_scale[i]
# 把当前 tile 的 softmax 权重与 KV value 相乘,累加到输出上。
T.gemm(acc_s_cast, KV_shared, acc_o, policy=T.GemmWarpPolicy.FullCol)
online softmax第二次循环+拷贝
和no split区别在于这里只是中间结果,后面还有多个seqplen split块合并,为了合并指数方便,这里转成指数和的对数logsum[i] = T.log2(logsum[i]) + scores_max[i] * scale
T.copy(logsum, glse[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, bz])指数和保存到前面申请的全局内存的对应位置,也就是bz
T.copy(O_shared, Output_partial[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, bz, :])eqkTve^{qk^T}veqkTv结果也保存到前面申请的全局内存的split对应下标
py
# 第一阶段结束后,acc_o 是当前 split 范围内的局部 value 加权和。
# 再除以该 split 内部的分母,得到局部归一化输出。
for i, j in T.Parallel(block_H, dim):
acc_o[i, j] /= logsum[i]
# 第二阶段需要的不是普通分母,而是 log-sum-exp 形式。
# 这里把分母重新变回对数域,便于跨 split 稳定归约。
for i in T.Parallel(block_H):
logsum[i] = T.log2(logsum[i]) + scores_max[i] * scale
T.copy(logsum, glse[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, bz])
T.copy(acc_o, O_shared)
T.copy(O_shared, Output_partial[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, bz, :])
第二个分块循环
为了合并前面的seqlen分块结果,还需要一次分块循环,这种合并操作,一般都对合并后维度分块,比如这里合并后就是(batch,head,dim),没有num_split了,就只对batch head两个维度分块,然后每个线程块内部串行枚举所有split,累加。
py
# 第二阶段负责跨 split 合并。
# 每个 CTA 处理一个输出 head 和一个 batch 元素。
with T.Kernel(heads, batch, threads=128) as (hid, bz):
申请块内存
po_local 保存每个split窗口的注意力输出,输出维度是dim。o_accum_local 累加最终输出,多个split的结果,乘上缩放因子,累加到这个位置,类型是fp32,不是中间变量的fp16,提高累加精度
lse_local_split 单变量,保存一个split块的指数和,lse_logsum_local 保存所有split块的指数和。
lse_max_local 维护所有split块的指数最值,用于缩放。scale_local 计算一个块的缩放比例。
py
# po_local 暂存某个 split 的局部输出向量。
po_local = T.alloc_fragment([dim], dtype)
# o_accum_local 用 fp32 累积最终输出。
o_accum_local = T.alloc_fragment([dim], accum_dtype)
# lse_local_split 是某个 split 的局部 log-sum-exp。
lse_local_split = T.alloc_var(accum_dtype)
# lse_logsum_local 是所有 split 合并后的总 log-sum-exp。
lse_logsum_local = T.alloc_var(accum_dtype)
# lse_max_local 是所有 split 中最大的 log-sum-exp。
# 它用于归约时的数值稳定。
lse_max_local = T.alloc_var(accum_dtype)
# scale_local 表示某个 split 对最终输出的贡献权重。
scale_local = T.alloc_var(accum_dtype)
T.clear(lse_logsum_local)
T.clear(o_accum_local)
求块指数和最大值
我们现在相当于有多个split块的结果,要对他们先根据max缩放,然后求出总和sum,再用总和进行softmax的归一化。
第一步先求出各个块元素和的最大值
py
# 先找到所有 split 中最大的 log-sum-exp。
# 后面所有 split 都基于这个基准做 exp2,避免数值不稳定。
lse_max_local = -T.infinity(accum_dtype)
for k in T.serial(num_split):
lse_max_local = T.max(lse_max_local, glse[bz, hid, k])
缩放各块的指数和并累加
利用最大值缩放各个指数和,累加到lse_logsum_local,累加是要经过指数T.exp2(lse_local_split - lse_max_local),因为保存的glse是对数
lse_logsum_local = T.log2(lse_logsum_local) + lse_max_local最后再变成对数。
py
# 重建所有 split 的分母贡献,并求出总的 log-sum-exp。
for k in T.Pipelined(num_split, num_stages=1):
lse_local_split = glse[bz, hid, k]
lse_logsum_local += T.exp2(lse_local_split - lse_max_local)
lse_logsum_local = T.log2(lse_logsum_local) + lse_max_local
利用全局指数和归一化
当前块的具体注意力拷贝出来po_local[i] = Output_partial[bz, hid, k, i]。
当前块的指数和lse_local_split = glse[bz, hid, k],和前面求的全局指数和一块,确定这块的缩放比例scale_local = T.exp2(lse_local_split - lse_logsum_local)
考虑缩放比例,累加到全局注意力输出o_accum_local[i] += po_local[i] * scale_local
Output[bz, hid, i] = o_accum_local[i]最后把结果累加块写回全局内存Output
py
# 每个 split 的局部输出已经在其内部完成了归一化。
# 这里再按 exp(lse_k - lse_total) 做加权求和,就得到全局结果。
for k in T.serial(num_split):
for i in T.Parallel(dim):
po_local[i] = Output_partial[bz, hid, k, i]
lse_local_split = glse[bz, hid, k]
scale_local = T.exp2(lse_local_split - lse_logsum_local)
for i in T.Parallel(dim):
o_accum_local[i] += po_local[i] * scale_local
# 把所有 split 归约后的结果写回最终输出。
for i in T.Parallel(dim):
Output[bz, hid, i] = o_accum_local[i]
完整实现
py
# split 版本会把 KV 序列分成 num_split 段。
# 第一阶段并行处理每一段,先得到局部 softmax 输出和对应的 log-sum-exp。
@T.prim_func
def main_split(
Q: T.Tensor([batch, heads, dim], dtype),
Q_pe: T.Tensor([batch, heads, pe_dim], dtype),
KV: T.Tensor([batch, seqlen_kv, kv_head_num, dim], dtype),
K_pe: T.Tensor([batch, seqlen_kv, kv_head_num, pe_dim], dtype),
Output: T.Tensor([batch, heads, dim], dtype),
):
# glse 保存每个 split 的 log-sum-exp。
# 第二阶段合并不同 split 时,需要依赖这个量重新构造全局 softmax 分母。
glse = T.alloc_global([batch, heads, num_split], dtype)
# Output_partial 保存每个 split 内部已经归一化后的局部输出。
# 第二阶段会用全局 softmax 权重把这些局部输出重新加权求和。
Output_partial = T.alloc_global([batch, heads, num_split, dim], dtype)
# 第一阶段的 grid 是 (batch, head_block, split_id)。
# 每个 CTA 负责一个 batch、一个 query head block、一个 KV split。
with T.Kernel(batch, heads // min(block_H, kv_group_num), num_split, threads=256) as (bid, hid, bz):
# Q_shared 缓存当前 CTA 负责的 query 内容部分。
Q_shared = T.alloc_shared([block_H, dim], dtype)
# S_shared 缓存 softmax 后的 score tile。
# 它会作为后续 value 侧 gemm 的输入。
S_shared = T.alloc_shared([block_H, block_N], dtype)
# Q_pe_shared 缓存 query 的位置编码部分。
Q_pe_shared = T.alloc_shared([block_H, pe_dim], dtype)
# KV_shared 缓存当前 block 的内容部分。
KV_shared = T.alloc_shared([block_N, dim], dtype)
# K_pe_shared 缓存当前 block 的位置编码部分。
K_pe_shared = T.alloc_shared([block_N, pe_dim], dtype)
# O_shared 用于把寄存器中的输出暂存后再写回全局内存。
O_shared = T.alloc_shared([block_H, dim], dtype)
# acc_s 是 score tile 的累加寄存器。
# 它承接"内容项 gemm"和"位置项 gemm"两次累加。
acc_s = T.alloc_fragment([block_H, block_N], accum_dtype)
# acc_s_cast 用于把 softmax 权重降回 fp16。
# 这样下一步做 value 侧 gemm 时更容易命中高吞吐路径。
acc_s_cast = T.alloc_fragment([block_H, block_N], dtype)
# acc_o 保存输出向量的累加结果。
# 它会随着每个 KV tile 的处理不断更新。
acc_o = T.alloc_fragment([block_H, dim], accum_dtype)
# scores_max 保存每一行当前见过的最大 score。
scores_max = T.alloc_fragment([block_H], accum_dtype)
# scores_max_prev 保存上一轮迭代的最大 score。
# 这样可以做 online softmax 的重缩放。
scores_max_prev = T.alloc_fragment([block_H], accum_dtype)
# scores_scale 是 old_max 切换到 new_max 后的重缩放因子。
scores_scale = T.alloc_fragment([block_H], accum_dtype)
# scores_sum 保存当前 tile 的 softmax 分子和。
scores_sum = T.alloc_fragment([block_H], accum_dtype)
# logsum 保存到当前 tile 为止累计出来的 softmax 分母。
logsum = T.alloc_fragment([block_H], accum_dtype)
# cur_kv_head 表示当前 head block 对应的 kv head。
# 因为这里限定 kv_head_num == 1,所以本质上总会映射到同一个 kv head。
cur_kv_head = hid // (kv_group_num // block_H)
# swizzle 用于优化线程块访问模式,尽量改善 shared memory 和全局内存行为。
T.use_swizzle(10)
# query 会反复与多个 KV tile 相乘。
# 先把对应的 Q 和 Q_pe 搬进 shared memory,可以减少重复读取。
T.copy(Q[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, :], Q_shared)
T.copy(Q_pe[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, :], Q_pe_shared)
# 初始时还没有累积任何输出。
T.fill(acc_o, 0)
# 初始分母设为 0。
T.fill(logsum, 0)
# 初始行最大值设为负无穷,方便后面逐 tile 更新。
T.fill(scores_max, -T.infinity(accum_dtype))
# 当前 split 只覆盖整段 KV 中的一部分长度。
# 这里继续按 block_N 做更细粒度分块,并以流水方式迭代。
loop_range = T.ceildiv((seqlen_kv // num_split), block_N)
for k in T.Pipelined(loop_range, num_stages=2):
kv_start = (seqlen_kv // num_split) * bz + k * block_N
kv_end = (seqlen_kv // num_split) * bz + (k + 1) * block_N
# 把当前 KV tile 的内容和位置编码一并读入 shared memory。
T.copy(KV[bid, kv_start:kv_end, cur_kv_head, :], KV_shared)
T.copy(K_pe[bid, kv_start:kv_end, cur_kv_head, :], K_pe_shared)
# 先清空 score 累加器。
T.clear(acc_s)
# 第一项是内容部分的点积。
T.gemm(Q_shared, KV_shared, acc_s, transpose_B=True, policy=T.GemmWarpPolicy.FullCol)
# 第二项是位置编码部分的点积。
# 两次 gemm 的结果相加后,数学上等价于先 concat 再做一次大矩阵乘。
T.gemm(Q_pe_shared, K_pe_shared, acc_s, transpose_B=True, policy=T.GemmWarpPolicy.FullCol)
# 下面进入 online softmax。
# 核心思想是边扫描 KV,边更新最大值和分母,而不是保存完整 score 矩阵。
T.copy(scores_max, scores_max_prev)
T.fill(scores_max, -T.infinity(accum_dtype))
T.reduce_max(acc_s, scores_max, dim=1, clear=False)
# 当前 tile 的最大值要与历史最大值合并,形成新的全局行最大值。
for i in T.Parallel(block_H):
scores_max[i] = T.max(scores_max[i], scores_max_prev[i])
# 历史分母和历史输出都需要按 max 的变化做重缩放。
for i in T.Parallel(block_H):
scores_scale[i] = T.exp2(scores_max_prev[i] * scale - scores_max[i] * scale)
# 把 score 转成 exp(score - max) 形式,得到当前 tile 的 softmax 分子。
for i, j in T.Parallel(block_H, block_N):
acc_s[i, j] = T.exp2(acc_s[i, j] * scale - scores_max[i] * scale)
# 累加每一行在当前 tile 上的分子和。
T.reduce_sum(acc_s, scores_sum, dim=1)
# 先写入 shared memory,再转成 fp16 fragment。
# 这是为了兼顾 softmax 阶段的精度和 gemm 阶段的吞吐。
T.copy(acc_s, S_shared)
T.copy(S_shared, acc_s_cast)
# 用重缩放后的旧分母加上当前 tile 的新分子和,得到新的分母。
for i in T.Parallel(block_H):
logsum[i] = logsum[i] * scores_scale[i] + scores_sum[i]
# 旧输出也必须同步按同样比例缩放。
# 这是 online softmax 保持数值正确的关键步骤。
for i, j in T.Parallel(block_H, dim):
acc_o[i, j] *= scores_scale[i]
# 把当前 tile 的 softmax 权重与 KV value 相乘,累加到输出上。
T.gemm(acc_s_cast, KV_shared, acc_o, policy=T.GemmWarpPolicy.FullCol)
# 第一阶段结束后,acc_o 是当前 split 范围内的局部 value 加权和。
# 再除以该 split 内部的分母,得到局部归一化输出。
for i, j in T.Parallel(block_H, dim):
acc_o[i, j] /= logsum[i]
# 第二阶段需要的不是普通分母,而是 log-sum-exp 形式。
# 这里把分母重新变回对数域,便于跨 split 稳定归约。
for i in T.Parallel(block_H):
logsum[i] = T.log2(logsum[i]) + scores_max[i] * scale
T.copy(logsum, glse[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, bz])
T.copy(acc_o, O_shared)
T.copy(O_shared, Output_partial[bid, hid * VALID_BLOCK_H : (hid + 1) * VALID_BLOCK_H, bz, :])
# 第二阶段负责跨 split 合并。
# 每个 CTA 处理一个输出 head 和一个 batch 元素。
with T.Kernel(heads, batch, threads=128) as (hid, bz):
# po_local 暂存某个 split 的局部输出向量。
po_local = T.alloc_fragment([dim], dtype)
# o_accum_local 用 fp32 累积最终输出。
o_accum_local = T.alloc_fragment([dim], accum_dtype)
# lse_local_split 是某个 split 的局部 log-sum-exp。
lse_local_split = T.alloc_var(accum_dtype)
# lse_logsum_local 是所有 split 合并后的总 log-sum-exp。
lse_logsum_local = T.alloc_var(accum_dtype)
# lse_max_local 是所有 split 中最大的 log-sum-exp。
# 它用于归约时的数值稳定。
lse_max_local = T.alloc_var(accum_dtype)
# scale_local 表示某个 split 对最终输出的贡献权重。
scale_local = T.alloc_var(accum_dtype)
T.clear(lse_logsum_local)
T.clear(o_accum_local)
# 先找到所有 split 中最大的 log-sum-exp。
# 后面所有 split 都基于这个基准做 exp2,避免数值不稳定。
lse_max_local = -T.infinity(accum_dtype)
for k in T.serial(num_split):
lse_max_local = T.max(lse_max_local, glse[bz, hid, k])
# 重建所有 split 的分母贡献,并求出总的 log-sum-exp。
for k in T.Pipelined(num_split, num_stages=1):
lse_local_split = glse[bz, hid, k]
lse_logsum_local += T.exp2(lse_local_split - lse_max_local)
lse_logsum_local = T.log2(lse_logsum_local) + lse_max_local
# 每个 split 的局部输出已经在其内部完成了归一化。
# 这里再按 exp(lse_k - lse_total) 做加权求和,就得到全局结果。
for k in T.serial(num_split):
for i in T.Parallel(dim):
po_local[i] = Output_partial[bz, hid, k, i]
lse_local_split = glse[bz, hid, k]
scale_local = T.exp2(lse_local_split - lse_logsum_local)
for i in T.Parallel(dim):
o_accum_local[i] += po_local[i] * scale_local
# 把所有 split 归约后的结果写回最终输出。
for i in T.Parallel(dim):
Output[bz, hid, i] = o_accum_local[i]