这几周处于个人调整的 GAP 期,逐渐整理过去的学习笔记。本文是我 2025 年 5 月份左右记录的散乱学习笔记,国庆假期做了整理,分享给初学者。
1 简介
注意力计算是大模型推理的关键环节,其核心矩阵运算包括三部分:计算 Q、K、V 矩阵、计算注意力、计算线性变换。本文将介绍单机、张量并行两种场景下的注意力计算过程,重点突出矩阵的变换,忽略位置编码。
2 单机注意力计算
2.1 计算流程图

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 计算流程图

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 矩阵

可以看到权重矩阵部分按列拆分给两个 tp worker。
(2)注意力结果的输出投影和 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 推荐学习资料
推荐新手可以进一步阅读以下文章,加深学习。
- Transformer 模型结构详解及代码实现
- 图解 Transformer
- Transformer 模型详解(图解最完整版)
- 基于 transformers 的自然语言处理(NLP)入门
- deepseek 技术解读(1)-彻底理解 MLA(Multi-Head Latent Attention)
- 构建 NLP 应用的架构演进之路
浙公网安备 33010602011771号