MLA 算子解析(Tilelang/Torch)

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]