这几周处于个人调整的 GAP 期,逐渐整理过去的学习笔记。本文是我 2025 年 5 月份左右记录的散乱学习笔记,国庆假期做了整理,分享给初学者。

1 简介

注意力计算是大模型推理的关键环节,其核心矩阵运算包括三部分:计算 Q、K、V 矩阵、计算注意力、计算线性变换。本文将介绍单机、张量并行两种场景下的注意力计算过程,重点突出矩阵的变换,忽略位置编码。

2 单机注意力计算

2.1 计算流程图

图 1|Transformer 单机注意力计算的整体流程

2.2 计算过程详解

2.2.1 关键参数说明

  • 一批有多少序列:batch_size
  • 一个序列的 token 个数:seq_len
  • 每个 token 的维度:d_model、hidden_size
  • 有多少头:num_heads
  • 每个头有多少维度:d_tensor = d_model / num_heads

2.2.2 计算过程

(1)【输入】用户输入的 prompt 经过 tokenizer 后转变为输入矩阵:[batch_size, seq_len, d_model]

(2)【权重】权重矩阵(通过模型训练得到,将所有 HEAD 的权重横向拼接到一起,方便单机计算):

  • Q 线性变换:w_q = [d_model, d_model]
  • K 线性变换:w_k = [d_model, d_model]
  • V 线性变换:w_v = [d_model, d_model]

注意:权重矩阵是做同维数投影的,第 0 维需要和输入 X 做运算,所以必须是 d_model。另外,第 1 维(多个 HEAD 合并后)通常也是 d_model。

(3)【qkv】通过线性变换得到 q、k、v 矩阵:输入矩阵 [batch_size, seq_len, d_model] * 权重矩阵 [d_model, d_model] = [batch_size, seq_len, d_model]

(4)【多头】使用 tensor.view 函数将 q、k、v 三个张量的最后一维拆为 num_heads 个二维矩阵,也就是按列拆。[batch_size, seq_len, d_model] ==> [batch_size, seq_len, num_heads, d_tensor],然后交换最后两维:[batch_size, seq_len, num_heads, d_tensor] ==> [batch_size, num_heads, seq_len, d_tensor],因为矩阵计算默认使用最后两维,而我们每一个 HEAD 需要计算的是 [seq_len, d_tensor],所以需要交换。

(5)【attention】执行 attention 计算,计算过程的矩阵变换(省略非矩阵计算部分):

  • q 和 k 转置的运算:[batch_size, num_heads, seq_len, d_tensor] * [batch_size, num_heads, d_tensor, seq_len] = [batch_size, num_heads, seq_len, seq_len]
  • *v 的运算:[batch_size, num_heads, seq_len, seq_len] * [batch_size, num_heads, seq_len, d_tensor] = [batch_size, num_heads, seq_len, d_tensor]

最后结果形状为:[batch_size, num_heads, seq_len, d_tensor]

(6)【concat】调换 1、2 维之后,将多个头合并转换:[batch_size, num_heads, seq_len, d_tensor] ==> [batch_size, seq_len, d_model],这个形状又和输入一样。

(7)【投影】执行线性变换投影:[batch_size, seq_len, d_model] * [d_model, d_model] = [batch_size, seq_len, d_model]

注意:

  • q / k / v 的权重矩阵都将各自的多个权重头横向拼接到一起,算出来之后再拆分为多头
  • attention 计算时必须拆分为多个头独立计算,不能直接拿合并 HEAD 的 d_model 去计算注意力,虽然结果维度一样,但是计算过程混在一起,是错误的

2.3 注意力计算的复杂度推导

推导一个长度为 S 的序列的计算复杂度,假设计算的基本单位为两个数字的乘法,hidden state 的维度用 H 表示。

(1)q、k、v 矩阵计算成本:

  • 矩阵变换:[S, H] * [H, H] = [S, H]
  • 复杂度算法:结果元素个数 * 算出每个元素要做的乘法个数 = S*H * H

(2)注意力计算成本:

  • 矩阵变换:[S, H] * [H, S] * [S, H] = [S, S] * [S, H] = [S, H]
  • 复杂度算法:S*S * H + S*H * S = 2*S*S*H

(3)汇总:S*H^2 + 2*S^2*H,对同一个模型的不同序列来说,H 是常量,因此:O(S) = S^2,注意力计算的复杂度和序列长度的平方成正比。

3 SGLang tp 注意力计算(MHA模型)

本章代码主要参考 qwen3.py 的 Qwen3Attention 实现,包含了三部分:qkv 矩阵计算(QKVParallelLinear)、注意力计算(RadixAttention)、输出投影矩阵计算(RowParallelLinear)。

其推理入口代码:

def forward(
    self,
    positions: torch.Tensor,
    hidden_states: torch.Tensor,
    forward_batch: ForwardBatch,
) -> torch.Tensor:
    qkv, _ = self.qkv_proj(hidden_states)
    q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
    q, k = self._apply_qk_norm(q, k)
    q, k = self.rotary_emb(positions, q, k)
    attn_output = self.attn(q, k, v, forward_batch)
    output, _ = self.o_proj(attn_output)
    return output

3.1 计算流程图

图 2|TP 模式下注意力计算的流程

3.2 计算过程详解

主要代码在 models 目录和 linear.py。

(1)【加载权重】进程初始化时加载权重矩阵。根据 tp_size、tp_rank 决定本 tp worker 需要负责哪些 HEAD,只加载这部分的权重(部分列)。最终看到的形状:[d_model, q+k+v 在本机的头权重]。

实际执行时有两种模式,一种是所有权重都加载,然后再 narrow(提取张量的子视图),最后覆盖掉完整的权重变量,确保没有被引用的显存能释放;另一种是磁盘文件已按照 tp 分好,每个 tp 只加载自己的。

(2)【输入矩阵】[batch_size, seq_len, d_model],所有 tp worker 的输入矩阵都是一样的。

(3)【q/k/v 矩阵】这里为了加速,将 q/k/v 矩阵的计算合并为一次矩阵计算:输入矩阵 [batch_size, seq_len, d_model] * 权重矩阵 [d_model, q+k+v 在本机的头权重]。在第 2 章中,我们看到多头是在 q/k/v 矩阵算完之后再拆分多头,这里我们提前在权重矩阵加载时就拆分多头。计算完后,再拆分出 q、k、v 矩阵,这里不会显式的提出 HEAD 维,而是都将本 tp 负责的多个 HEAD 对应的值汇总在最后一维,譬如:[batch_size, seq_len, num_heads * head_dim]。注:在 GQA 模式下,q 的 HEAD 数和 k/v 的 HEAD 数是不一样的,为便于理解,我们忽略这部分差异。

(4)【attention】每个 tp 都执行 attention 计算,同第 2 章介绍的类似。

(5)【投影】每个 tp 都执行投影:[batch_size, seq_len, num_heads * head_dim] * [num_heads * head_dim, d_model],将维度变回 [batch_size, seq_len, d_model]。这里的投影权重矩阵也是在启动时每个 tp 只加载部分权重矩阵(部分行),这里是按照行拆分的,所以 sglang 里的类名字叫做 RowParallelLinear。

(6)【all_reduce】每一个 tp 上都执行 all_reduce 汇总结果(加和),代码在 RowParallelLinear 类的 forward 里。

为什么最后是 all_reduce? q/k/v 矩阵计算的时候,采用多头拆分,这是横向(按列)拆分,一个 tp 算出来的结果是整体结果的一部分列。再执行投影矩阵计算,投影的权重矩阵是按行拆分(这样矩阵乘法中间的维度才能对得上),此时一个 tp 算出来的形状和最终结果一致,但每个矩阵元素的值只是最终值的一部分,需要加和。这和单机模式的计算不同,单机模式下,先横向拼接为完整的注意力结果矩阵,然后再做投影计算。因此最后的投影部分可以总结为两种模式:单机模式先拼接后计算,或者先计算后加和。

在 tp 模式下,一个 tp 算出来的形状和最终结果一致,但每个矩阵元素的值只是最终值的一部分,了解到这个信息之后,可以较快的理解 scatter 模式(只给 MLP 层送部分 token 的注意力结果做激活)下的 attn_tp_reduce_scatter 操作。

3.3 计算过程的矩阵变换图展

(1)求 Q/K/V 矩阵

图 3|TP 模式下求 Q\K\V 矩阵:权重按列切分

可以看到权重矩阵部分按列拆分给两个 tp worker。

(2)注意力结果的输出投影和 all_reduce

图 4|TP 模式下的输出投影与 all_reduce

可以看到每一个 tp worker 得到的矩阵形状都是最终形状,但元素的值是部分值,需要按位置相加。

3.4 为什么 TP 按照 HEAD 个数拆分?

在配置 tp size 的时候,tp_size 和 HEAD 个数必须是有整除关系,譬如 qwen3.py 的 Qwen3Attention 代码:

if self.total_num_kv_heads >= self.tp_size:
    # Number of KV heads is greater than TP size, so we partition
    # the KV heads across multiple tensor parallel GPUs.
    assert self.total_num_kv_heads % self.tp_size == 0
else:
    # Number of KV heads is less than TP size, so we replicate
    # the KV heads across multiple tensor parallel GPUs.
    assert self.tp_size % self.total_num_kv_heads == 0

有整除关系是为了后面矩阵按 TP 切分时,可以确保每一个 TP 上的 HEAD 是完整的,为什么需要完整的 HEAD 呢,或者说为什么 TP 需要按照 HEAD 拆分,而不是按列维度自由拆分呢?

这是注意力计算的多头特性和张量并行的结合,如果不按照 HEAD 拆分,就需要引入卡间通信,得不偿失。参考本文第 2 章「单机注意力计算」的介绍,注意力计算的矩阵乘法分为两步,第一步是 q 乘以 k 的转置,第二步是前述结果乘以 v。假设 tp=2,并且 q、k、v 矩阵都只有半个 HEAD 维度,那么在每一个 TP 上,q 乘以 k 转置得到的矩阵元素个数是完整的 [seq_len, seq_len](忽略其他维度),但是每一个元素的值只有一半,需要两个 tp rank 的 q 乘以 k 转置的结果加和才能得到完整的结果(即 AllReduce),这就需要引入卡间通信。

4 推荐学习资料

推荐新手可以进一步阅读以下文章,加深学习。

本文所在:https://www.cnblogs.com/cswuyg/p/23217280

posted on 2026-10-07 22:13  -银光-  阅读(10)  评论(0)    收藏  举报