深度学习
Transformer 矩阵与维度完整讲解(整合版)
从 Token 输入、Q/K/V、多头注意力到 Transformer Block 与 KV Cache,系统追踪矩阵运算和维度变化。
主题:从 Token ID、Embedding、Q/K/V、缩放点积注意力和多头注意力,一直推导到 Transformer Block、词表 logits、Prefill、Decode 与 KV Cache。
适用范围:以 Decoder-only Transformer(例如 GPT 类模型)为主,同时说明标准 Transformer 与现代实现中的常见差异。
阅读目标:不仅记住公式,还能解释每一步在做什么、矩阵为什么可以相乘、形状怎样变化。
一、先看完整主线
Token IDs
[B,T]
↓ Embedding
输入 X
[B,T,D]
↓ Q、K、V 线性投影
↓ 拆成 H 个注意力头
Q、K、V
[B,H,T,Dh]
↓ Q × Kᵀ
注意力分数
[B,H,T,T]
↓ 除以 √Dh
↓ Mask
↓ Softmax
注意力权重
[B,H,T,T]
↓ 权重 × V
各注意力头的上下文向量
[B,H,T,Dh]
↓ 拼接 H 个头
↓ 输出投影 WO
多头注意力输出
[B,T,D]
↓ 残差、归一化、FFN
Transformer Block 输出
[B,T,D]
↓ 重复 N 层
↓ 词表输出层
logits
[B,T,N_vocab]
最核心的形状变化是:
[B,T,D]
→ [B,H,T,Dh]
→ [B,H,T,T]
→ [B,H,T,Dh]
→ [B,T,D]
其中:
D = H × Dh
先不用急着理解每个符号。下面会从输入开始逐步解释。
二、统一符号
| 符号 | 含义 |
|---|---|
B | Batch Size,一次处理的样本数 |
T | Sequence Length,一条输入中的 Token 数 |
D | Model Dimension,每个 Token 的主向量长度 |
H | Attention Heads,注意力头数量 |
Dh | Head Dimension,每个注意力头的向量长度,标准 MHA 中通常 Dh=D/H |
Dff | FFN 的中间层维度 |
N_vocab | 词表中的 Token 总数 |
N | Transformer Block 的层数 |
X | 进入注意力模块的输入 |
Q | Query,查询向量 |
K | Key,匹配向量 |
V | Value,内容向量 |
A | Attention Weights,注意力权重 |
为了避免混淆:
V表示 Value 矩阵;N_vocab表示词表大小;WQ、WK、WV、WO是训练得到的参数;Q、K、V、A是模型处理当前输入时临时计算出来的结果。
三、从 Token ID 到输入矩阵
假设:
B = 2:一次输入两条文本
T = 4:每条文本整理成 4 个 Token
D = 8:每个 Token 用 8 个数字表示
分词后,Token ID 的形状是:
token_ids: [2,4]
如果词表大小为 10,000,Embedding 参数表的形状是:
E: [10000,8]
Embedding 不是把 [2,4] 与 [10000,8] 做普通矩阵乘法,而是把每个 Token ID 当作行号,去参数表中查询对应向量:
X_token = E[token_ids]
查询后得到:
X_token: [2,4,8]
含义是:
2 条输入
× 每条 4 个 Token
× 每个 Token 8 维
3.1 位置信息怎样加入
仅有 Token Embedding 时,相同 Token 在不同位置会先得到相同的词向量,因此模型还需要位置信息。
如果使用可学习位置向量或原始 Transformer 的位置编码,可以简化为:
X = X_token + X_position
两者形状相同:
[B,T,D] + [B,T,D] = [B,T,D]
如果模型使用 RoPE,位置信息主要在后面作用于 Q、K,而不是直接加到 X 上。无论采用哪种方式,主形状仍然保持:
X: [B,T,D]
四、矩阵乘法和批量矩阵乘法
4.1 最基本的形状规则
如果:
A: [m,n]
B: [n,p]
那么:
A × B: [m,p]
中间两个维度 n 必须相同。
例如:
[3,2] × [2,4] = [3,4]
不能直接进行:
[3,2] × [3,4]
因为中间维度 2 和 3 不相等。
4.2 Batch 和 Head 维度怎样参与计算
看到:
Q: [B,H,T,Dh]
Kᵀ: [B,H,Dh,T]
不要把四个维度全部混在一起相乘。真正发生的是:
对每个 Batch、每个 Head,分别执行最后两个维度的矩阵乘法。
也就是一共执行 B×H 组:
[T,Dh] × [Dh,T] = [T,T]
再把所有结果组织成:
[B,H,T,T]
B 和 H 是批量维度,最后两个维度才是每次矩阵乘法真正使用的行和列。
五、从输入 X 生成 Q、K、V
进入注意力模块的输入是:
X: [B,T,D]
从单个头的概念看,三组参数可以写成:
WQ: [D,Dh]
WK: [D,Dh]
WV: [D,Dh]
计算:
Q = X × WQ
K = X × WK
V = X × WV
忽略 Batch 后:
X [T,D] × WQ [D,Dh] = Q [T,Dh]
X [T,D] × WK [D,Dh] = K [T,Dh]
X [T,D] × WV [D,Dh] = V [T,Dh]
5.1 为什么同一个 X 要生成三种向量
同一个 Token 在注意力中需要扮演三种角色:
Q:当前 Token 想寻找什么信息;K:当前 Token 能用什么特征参与匹配;V:如果当前 Token 被关注,它实际提供什么内容。
可以把一次注意力理解成:
- 用当前 Token 的 Q 去匹配所有 Token 的 K;
- 根据匹配结果得到注意力权重;
- 用这些权重组合所有 Token 的 V。
WQ、WK、WV 是模型训练得到的三组不同参数,所以即使 Q、K、V 都来自同一个 X,它们通常也不相等。
5.2 参数与中间结果不要混淆
| 类型 | 示例 | 推理时是否随输入变化 |
|---|---|---|
| 模型参数 | WQ、WK、WV、WO | 同一个已训练模型中保持不变 |
| 当前输入的中间结果 | Q、K、V、分数、注意力权重 | 会随输入变化 |
注意力权重不是模型训练后固定保存的一张表。输入句子改变,Q、K 和注意力权重也会改变。
5.3 真实多头实现中的合并写法
实际程序往往不会为每个头单独调用 H 次线性层,而是先使用大参数矩阵:
WQ、WK、WV: [D,D]
得到:
Q、K、V: [B,T,D]
再把最后一个 D 拆成:
D = H × Dh
还有一些实现会把三次投影合并为一次:
WQKV: [D,3D]
X × WQKV → [B,T,3D]
随后再切分出 Q、K、V。这只是工程实现的合并,数学角色没有改变。
六、完整手算一次单头自注意力
为了看清每个数字的来源,设定:
B = 1
T = 3
D = 2
H = 1
Dh = 2
暂时省略 Batch 维度。三个 Token 的输入向量为:
X = [
[1,0],
[0,1],
[1,1]
]
形状:
X: [3,2]
为了展示 Q、K、V 确实来自不同投影,使用三张不同参数矩阵:
WQ = [
[1,0],
[0,1]
]
WK = [
[1,0],
[1,1]
]
WV = [
[1,1],
[0,1]
]
三张参数矩阵的形状都是 [2,2]。这些数值只用于演示,不是真实模型参数。
6.1 计算 Q
Q = X × WQ
因为 WQ 是单位矩阵:
Q = [
[1,0],
[0,1],
[1,1]
]
形状:
[3,2] × [2,2] = [3,2]
6.2 计算 K
K = X × WK
逐行计算:
[1,0] × WK = [1,0]
[0,1] × WK = [1,1]
[1,1] × WK = [2,1]
所以:
K = [
[1,0],
[1,1],
[2,1]
]
6.3 计算 V
V = X × WV
逐行计算:
[1,0] × WV = [1,1]
[0,1] × WV = [0,1]
[1,1] × WV = [1,2]
所以:
V = [
[1,1],
[0,1],
[1,2]
]
现在可以直接看到:
Q ≠ K ≠ V
6.4 用 Q 和 K 计算匹配分数
原始分数为:
Q × Kᵀ
现在:
Q: [3,2]
Kᵀ: [2,3]
因此:
[3,2] × [2,3] = [3,3]
K 的转置是:
Kᵀ = [
[1,1,2],
[0,1,1]
]
矩阵乘法结果:
QKᵀ = [
[1,1,2],
[0,1,1],
[1,2,3]
]
第一行第三列来自:
q1 · k3
= [1,0] · [2,1]
= 1×2 + 0×1
= 2
分数矩阵的阅读方式是:
| 当前 Query | Key 1 | Key 2 | Key 3 |
|---|---|---|---|
| Token 1 | 1 | 1 | 2 |
| Token 2 | 0 | 1 | 1 |
| Token 3 | 1 | 2 | 3 |
- 每一行对应一个当前 Token 的 Query;
- 每一列对应一个被匹配 Token 的 Key;
- 第
i行第j列表示 Tokeni的 Query 与 Tokenj的 Key 的匹配分数。
6.5 为什么要除以 √Dh
缩放后的分数是:
S = QKᵀ / √Dh
这里:
Dh = 2
√Dh = √2 ≈ 1.414
所以:
S ≈ [
[0.707,0.707,1.414],
[0, 0.707,0.707],
[0.707,1.414,2.121]
]
直观上,Dh 越大,点积需要累加的项越多,分数的绝对值容易变大。大数进入 Softmax 后,概率会过早集中在少数位置,梯度也可能变得不稳定。
更具体地说,如果 Q 和 K 的各维近似独立、均值约为 0、方差约为 1,那么:
Var(q · k) ≈ Dh
Std(q · k) ≈ √Dh
除以 √Dh 后,可以把分数的尺度控制在更稳定的范围。
6.6 Softmax 把分数变成权重
Softmax 对每个 Query 所在的行分别计算:
softmax(s_i) = exp(s_i) / Σ exp(s_j)
第一行:
[0.707,0.707,1.414]
取指数:
[2.028,2.028,4.113]
归一化后约为:
[0.248,0.248,0.503]
三行的注意力权重约为:
A ≈ [
[0.248,0.248,0.503],
[0.198,0.401,0.401],
[0.140,0.284,0.576]
]
每一行的总和约为 1。它表示一个 Query 应该按什么比例读取各个 Value。
实际程序通常先让每一行减去该行最大值,再计算指数:
softmax(s_i)
= exp(s_i-max(s)) / Σ exp(s_j-max(s))
这样可以减少指数溢出,数学结果不变。
6.7 用权重组合 V
O = A × V
形状:
A: [3,3]
V: [3,2]
[3,3] × [3,2] = [3,2]
第一行的计算:
0.248×[1,1]
+ 0.248×[0,1]
+ 0.503×[1,2]
≈ [0.752,1.503]
完整输出约为:
O ≈ [
[0.752,1.503],
[0.599,1.401],
[0.716,1.576]
]
输入 X 中,每个 Token 原本只有自己的向量;输出 O 中,每个 Token 的新向量已经按注意力权重混入了其他 Token 的 Value 信息。
这就是自注意力的核心:
Attention(Q,K,V)
= softmax(QKᵀ / √Dh) V
七、Mask 为什么在 Softmax 之前加入
7.1 GPT 的因果 Mask
Decoder-only 模型生成当前位置时不能读取未来 Token,因此需要因果 Mask。
对于 T=3:
M = [
[0,-∞,-∞],
[0, 0,-∞],
[0, 0, 0]
]
完整公式变成:
Attention(Q,K,V)
= softmax(QKᵀ / √Dh + M) V
把 Mask 加到前面的缩放分数上:
S + M ≈ [
[0.707,-∞, -∞],
[0, 0.707,-∞],
[0.707, 1.414,2.121]
]
因为:
exp(-∞) = 0
Softmax 后:
A_causal ≈ [
[1, 0, 0],
[0.330,0.670,0],
[0.140,0.284,0.576]
]
含义:
- Token 1 只能读取 Token 1;
- Token 2 可以读取 Token 1 和 Token 2;
- Token 3 可以读取前三个 Token。
如果先做 Softmax 再加 Mask,被禁止位置已经参与了概率归一化,剩余权重之和也会被破坏。因此 Mask 必须在 Softmax 之前作用于分数。
7.2 Padding Mask
同一个 Batch 中的文本长度可能不同,较短文本会补 Padding Token。Padding Mask 用于阻止模型读取这些补齐位置。
在多头计算中,概念形状常写为:
padding_mask: [B,1,1,T]
它主要限制 Key 方向上哪些位置允许被读取。
7.3 Mask 的广播
因果 Mask 的概念形状可以写成:
causal_mask: [1,1,T,T]
注意力分数的形状是:
scores: [B,H,T,T]
Padding Mask 和 Causal Mask 可以通过广播作用到每个 Batch、每个 Head 的分数矩阵。
不同框架对布尔值含义、数值形式和存储形状的约定可能不同,使用具体 API 时要看文档;它们的最终目标都是让禁止位置在 Softmax 后获得 0 权重。
八、从单头扩展到多头注意力
设定:
B = 2
T = 4
D = 8
H = 2
Dh = D/H = 4
输入:
X: [2,4,8]
8.1 Q、K、V 线性投影
标准多头注意力常使用三张合并后的大参数矩阵:
WQ、WK、WV: [8,8]
所以:
Q = X × WQ: [2,4,8]
K = X × WK: [2,4,8]
V = X × WV: [2,4,8]
最后一个 8 包含:
H × Dh = 2 × 4
8.2 拆分注意力头
先重塑:
[B,T,D]
→ [B,T,H,Dh]
[2,4,8]
→ [2,4,2,4]
再交换维度顺序:
[B,T,H,Dh]
→ [B,H,T,Dh]
[2,4,2,4]
→ [2,2,4,4]
于是:
Q、K、V: [2,2,4,4]
含义是:
2 条样本
× 每条样本 2 个注意力头
× 每个头 4 个 Token
× 每个 Token 在该头中用 4 个数字表示
8.3 每个头分别生成 T×T 关系表
Q: [B,H,T,Dh]
Kᵀ: [B,H,Dh,T]
最后两个维度相乘:
[B,H,T,Dh] × [B,H,Dh,T]
→ [B,H,T,T]
代入示例:
[2,2,4,4] × [2,2,4,4]
→ [2,2,4,4]
这个例子中 T=4、Dh=4,所以 K 转置前后的数字看起来一样,但语义并不一样:
- K 转置前最后两维是
[T,Dh]; - K 转置后最后两维是
[Dh,T]。
分数张量表示:
2 条样本
× 每条 2 个注意力头
× 每个头一张 4×4 关系表
8.4 每个头用权重组合 V
经过缩放、Mask 和 Softmax 后:
A: [B,H,T,T]
V: [B,H,T,Dh]
计算:
[B,H,T,T] × [B,H,T,Dh]
→ [B,H,T,Dh]
示例输出:
Z: [2,2,4,4]
每个头都为每个 Token 生成一个 Dh 维上下文向量。
8.5 拼接多个头
先把维度顺序换回来:
[B,H,T,Dh]
→ [B,T,H,Dh]
再合并 H 和 Dh:
[B,T,H,Dh]
→ [B,T,H×Dh]
→ [B,T,D]
示例:
[2,2,4,4]
→ [2,4,2,4]
→ [2,4,8]
8.6 输出投影 WO
拼接后还会乘输出参数:
WO: [D,D]
计算:
O = Concat(head1,...,headH) × WO
示例:
[2,4,8] × [8,8]
→ [2,4,8]
WO 不只是为了修复形状,它还会学习怎样混合不同注意力头的输出。
8.7 多头为什么不是多个完整模型
多个注意力头:
- 处理同一批输入;
- 属于同一个 Transformer Block;
- 在不同参数投影形成的子空间中并行计算;
- 最终会被拼接并继续进入同一网络。
因此,多头注意力是一个模块内部的并行表示机制,不是同时运行 H 个完整大模型。
不同头可能学到不同关系,但不能预先规定某个头一定负责语法、另一个头一定负责指代。
8.8 标准 MHA、GQA 与 MQA
本文的基础推导使用标准 MHA:
Query 头数 = Key 头数 = Value 头数 = H
现代模型还可能使用:
- MQA:所有 Query 头共享一组 Key、Value;
- GQA:若干 Query 头共享一组 Key、Value。
它们主要用于减少 KV Cache 和推理开销。学习基础矩阵关系时先掌握标准 MHA,阅读具体模型代码时再确认 KV 头数。
九、完整 Transformer Block
注意力模块最终保持:
[B,T,D] → [B,T,D]
这是因为输出需要与原输入做残差相加,也因为多个 Block 需要连续堆叠。
9.1 残差连接
X: [B,T,D]
O: [B,T,D]
R = X + O
R: [B,T,D]
残差连接保留了一条原输入路径:
新表示 = 原表示 + 子模块学习到的增量信息
如果注意力结果仍是 [B,H,T,Dh],就不能直接与 X 相加,所以必须先拼接注意力头并投影回 [B,T,D]。
9.2 LayerNorm
对一个 Token 的 D 维向量:
x = [x1,x2,...,xD]
先计算均值与方差:
μ = (1/D) Σ xi
σ² = (1/D) Σ (xi-μ)²
再归一化,并加入可训练缩放和平移参数:
x_hat_i = (xi-μ) / √(σ²+ε)
y_i = γ_i × x_hat_i + β_i
整体形状不变:
[B,T,D] → [B,T,D]
LayerNorm 通常对每个 Token 自己的 D 个特征进行归一化,不依赖 Batch 中有多少样本。
9.3 RMSNorm
很多现代 Decoder-only 模型使用 RMSNorm:
RMS(x) = √((1/D) Σ xi² + ε)
y_i = γ_i × xi / RMS(x)
RMSNorm 通常不执行 LayerNorm 中的“减去均值”,也常不使用 β。输入输出形状仍然是:
[B,T,D] → [B,T,D]
9.4 FFN 前馈网络
注意力负责:
让不同 Token 位置之间交换信息。
FFN 负责:
对每个 Token 已经获得的特征继续加工。
基础 FFN:
FFN(x) = activation(xW1+b1)W2+b2
假设:
D = 8
Dff = 32
参数:
W1: [8,32]
W2: [32,8]
形状变化:
[2,4,8] × [8,32] → [2,4,32]
[2,4,32] × [32,8] → [2,4,8]
FFN 中间会扩大到 Dff,最终必须压回 D,才能做残差相加并进入下一层。
普通 position-wise FFN 对每个 Token 位置独立使用同一组参数,它本身不会在不同 Token 之间做注意力式的信息交换。
9.5 门控 FFN
一些现代模型使用 SwiGLU 等门控结构。简化形式:
gate = SiLU(x × W_gate)
up = x × W_up
hidden = gate ⊙ up
output = hidden × W_down
其中 ⊙ 表示对应位置逐元素相乘。
常见形状:
x: [B,T,D]
W_gate: [D,Dff]
W_up: [D,Dff]
gate: [B,T,Dff]
up: [B,T,Dff]
hidden: [B,T,Dff]
W_down: [Dff,D]
output: [B,T,D]
无论使用基础 FFN 还是门控 FFN,主形状都是:
[B,T,D] → [B,T,Dff] → [B,T,D]
9.6 Pre-Norm 与 Post-Norm
现代 Decoder-only 模型常见 Pre-Norm:
U = X + MHA(Norm(X))
Y = U + FFN(Norm(U))
原始 Transformer 论文中的经典结构更接近 Post-Norm:
U = Norm(X + MHA(X))
Y = Norm(U + FFN(U))
不同资料里的归一化位置不同,不表示其中一个一定错误。要先确认资料讲的是原始 Transformer,还是某个现代模型的具体结构。
9.7 一个完整 Pre-Norm Block 的形状
X = 输入 [B,T,D]
X1 = Norm(X) [B,T,D]
A = MultiHeadAttention(X1) [B,T,D]
U = X + A [B,T,D]
U1 = Norm(U) [B,T,D]
F1 = U1 × W1 [B,T,Dff]
F2 = activation(F1) [B,T,Dff]
F3 = F2 × W2 [B,T,D]
Y = U + F3 [B,T,D]
训练阶段还可能在注意力权重、注意力输出、FFN 中间层或残差分支中使用 Dropout;推理阶段通常关闭 Dropout。Dropout 不改变这里的主形状推导。
十、堆叠多层后怎样得到词表 logits
假设模型有 N 层:
X0: [B,T,D]
X1: [B,T,D]
X2: [B,T,D]
...
XN: [B,T,D]
输入输出形状相同,不代表内容没有变化。每一层都会使用自己的参数重新计算 Q、K、V、注意力和 FFN,Token 的语义表示会不断更新。
最后把每个 Token 的 D 维表示映射到整个词表:
W_vocab: [D,N_vocab]
如果:
XN: [2,4,8]
N_vocab = 10000
那么:
logits = XN × W_vocab
[2,4,8] × [8,10000]
→ [2,4,10000]
含义:
2 条样本
× 每条 4 个 Token 位置
× 每个位置对 10000 个候选 Token 打分
一些模型会让 W_vocab 与输入 Embedding 参数共享,称为 Weight Tying;是否共享要看具体模型架构。
10.1 训练时为什么保留所有位置
自回归语言模型训练时,每个有效位置都可以承担一次“预测下一个 Token”的任务,因此会保留:
logits: [B,T,N_vocab]
然后把预测与向后错一位的目标 Token 对齐,并在非 Padding 位置上计算损失。
10.2 生成时为什么只取最后一个位置
推理生成下一个 Token 时,当前只需要最后一个有效位置的预测:
last_logits = logits[:,-1,:]
last_logits: [B,N_vocab]
再在词表维度上进行选择。可以先转成概率:
P(next_token) = softmax(last_logits)
实际生成还可能使用 greedy、temperature、top-k、top-p 等策略。
10.3 两种 Softmax 不要混淆
| Softmax | 典型输入形状 | 归一化对象 | 作用 |
|---|---|---|---|
| 注意力 Softmax | [B,H,T,T] | 当前 Query 可读取的 Key 位置 | 决定从哪些上下文位置读取信息 |
| 词表 Softmax | [B,N_vocab] | 词表中的候选 Token | 得到下一个 Token 的概率分布 |
二者数学形式相同,但解决的问题完全不同。
十一、Prefill、Decode 与 KV Cache
前面的 [T,T] 推导解释了一次处理完整序列的情况。实际 GPT 推理通常分为 Prefill 和 Decode 两个阶段。
11.1 Prefill:第一次处理完整 Prompt
假设 Prompt 长度为 T。标准 MHA 中,每一层会计算:
Q: [B,H,T,Dh]
K: [B,H,T,Dh]
V: [B,H,T,Dh]
分数:
Q [B,H,T,Dh]
× Kᵀ [B,H,Dh,T]
→ [B,H,T,T]
Prefill 结束后,每一层都会保存 K 和 V:
K_cache: [B,H,T,Dh]
V_cache: [B,H,T,Dh]
11.2 Decode:每次生成一个新 Token
生成新 Token 时,只为新增位置计算:
q_new: [B,H,1,Dh]
k_new: [B,H,1,Dh]
v_new: [B,H,1,Dh]
把新 k、v 追加到缓存。假设追加后的上下文总长度为 T:
K_cache: [B,H,T,Dh]
V_cache: [B,H,T,Dh]
新 Query 与所有历史 Key 匹配:
q_new [B,H,1,Dh]
× K_cacheᵀ [B,H,Dh,T]
→ scores [B,H,1,T]
再组合所有历史 Value:
A [B,H,1,T]
× V_cache [B,H,T,Dh]
→ context [B,H,1,Dh]
因此,Decode 阶段不需要重新计算所有旧 Token 的 K 和 V,也不需要为旧位置重新生成完整的 [T,T] 输出;它只为新 Token 计算一行 [1,T] 的注意力。
11.3 为什么缓存 K 和 V,而不缓存旧 Q
生成新 Token 时,需要:
- 用新 Token 的 Query 去查询全部历史 Key;
- 根据权重读取全部历史 Value。
旧 Token 的 Query 是它们当时主动查询上下文时使用的,生成新 Token 时不再需要,所以 KV Cache 通常保存 K、V,不保存旧 Q。
11.4 KV Cache 为什么占显存
标准 MHA 中,每层 K、V 的元素数量近似为:
2 × B × H × T × Dh
如果模型有 N 层:
2 × N × B × H × T × Dh
还需要乘每个元素占用的字节数。FP16 或 BF16 通常每个元素约 2 字节。
使用 GQA 或 MQA 时,缓存中的 KV 头数少于 Query 头数,因此 KV Cache 会更小。真实推理服务还会有分页、对齐、调度等额外开销。
十二、计算量与内存主要花在哪里
12.1 注意力
注意力分数矩阵的形状是:
[B,H,T,T]
分数或权重的元素数量与:
B × H × T²
成正比。
QKᵀ 的主要计算量近似为:
B × H × T² × Dh
因为 H×Dh=D,也常写成:
O(B × T² × D)
所以标准全量注意力中,T 翻倍时,T² 项约变为原来的 4 倍。
12.2 FFN
FFN 的主要计算量近似与:
B × T × D × Dff
成正比。Dff 往往明显大于 D,所以 FFN 也会消耗大量参数和计算,不能只关注注意力矩阵。
12.3 参数量的粗略比较
忽略偏置,标准多头注意力的投影参数约为:
WQ、WK、WV、WO
≈ 4D²
基础 FFN 的参数约为:
W1、W2
≈ 2D×Dff
如果 Dff≈4D,基础 FFN 参数约为 8D²,可能比注意力投影参数更多。具体比例会因门控结构和模型配置而变化。
FlashAttention 等技术主要优化注意力的内存访问和计算组织,不会改变 Q、K、V 的基础语义和公式。
十三、一组尺寸串起全过程
设:
B = 2
T = 4
D = 8
H = 2
Dh = 4
Dff = 32
N_vocab = 10000
| 步骤 | 操作 | 输出形状 |
|---|---|---|
| Token IDs | 输入编号 | [2,4] |
| Embedding | 按 Token ID 查表 | [2,4,8] |
| Q、K、V 投影 | 分别乘 [8,8] | 各 [2,4,8] |
| 拆分多头 | reshape + transpose | 各 [2,2,4,4] |
QKᵀ | 最后两维相乘 | [2,2,4,4] |
| 缩放与 Mask | 除以 √4,再加 Mask | [2,2,4,4] |
| 注意力 Softmax | 按 Key 位置归一化 | [2,2,4,4] |
| 权重乘 V | A×V | [2,2,4,4] |
| 拼接多头 | 合并 H×Dh | [2,4,8] |
| 输出投影 | 乘 WO [8,8] | [2,4,8] |
| 残差相加 | X+Attention | [2,4,8] |
| FFN 扩大 | 乘 W1 [8,32] | [2,4,32] |
| FFN 压回 | 乘 W2 [32,8] | [2,4,8] |
| 再次残差 | U+FFN | [2,4,8] |
| 重复 N 层 | Transformer Blocks | [2,4,8] |
| 词表输出层 | 乘 [8,10000] | [2,4,10000] |
| 取最后位置 | logits[:,-1,:] | [2,10000] |
这里多个形状都恰好是 [2,2,4,4],但它们的语义不同:
- Q、K、V:最后两维是
[T,Dh]; - Kᵀ:最后两维是
[Dh,T]; - 注意力分数:最后两维是
[T_query,T_key]。
不能只看数字相同,就认为这些张量表示同一种东西。
十四、最容易出现的理解错误
错误 1:把 WQ 和 Q 当成同一个东西
WQ:模型参数,[D,Dh] 或合并后的 [D,D]
Q:当前输入经过 WQ 得到的中间结果,[B,H,T,Dh]
错误 2:认为 Q、K、V 是三段不同文本
自注意力中,Q、K、V 通常都由同一个输入 X 经过不同参数投影得到。交叉注意力中,它们的来源才可能不同。
错误 3:忘记 K 要转置
Q [T,Dh] × Kᵀ [Dh,T] = [T,T]
直接用 Q 乘 K,通常会因为中间维度不匹配而无法执行。
错误 4:把 D 和 Dh 混为一谈
D = 一个 Token 的完整模型维度
Dh = 一个注意力头的维度
D = H × Dh
错误 5:Softmax 做错维度
注意力 Softmax 通常沿最后一个 Key 位置维度归一化:
[B,H,T_query,T_key]
对于每个 Query,所有允许读取的 Key 的权重之和为 1。
错误 6:先 Softmax 再加 Mask
Mask 应在 Softmax 之前加到分数上,禁止位置才能在 Softmax 后得到 0 权重。
错误 7:把多个头直接相加
标准 MHA 通常先拼接:
[B,H,T,Dh] → [B,T,D]
再乘 WO,而不是把所有头直接逐元素相加。
错误 8:认为 FFN 会混合不同 Token
普通 FFN 对每个 Token 位置分别使用同一组参数。Token 之间的信息交换主要发生在注意力模块中。
错误 9:输入输出形状相同,所以内容没变
Transformer Block 保持 [B,T,D]→[B,T,D] 是为了残差连接和多层堆叠,内部数值和语义表示已经变化。
错误 10:把注意力权重当成下一个 Token 的概率
注意力权重:[B,H,T,T]
词表概率: [B,N_vocab]
二者的归一化对象和作用不同。
错误 11:认为 KV Cache 让生成速度不再受上下文长度影响
KV Cache 避免了重复计算旧 Token 的 K、V,但新 Query 仍要读取不断增长的历史 K、V。上下文越长,单步 Decode 的读取量和缓存占用仍会增加。
十五、最终复盘
15.1 一条公式
Attention(Q,K,V)
= softmax(QKᵀ / √Dh + M) V
15.2 一条形状主线
X [B,T,D]
→ Q、K、V [B,H,T,Dh]
→ Scores [B,H,T,T]
→ Context [B,H,T,Dh]
→ Concat [B,T,D]
→ FFN [B,T,Dff] → [B,T,D]
→ Block Output [B,T,D]
→ logits [B,T,N_vocab]
15.3 一段完整口述
Token ID 先通过 Embedding 查表变成[B,T,D]的输入 X。X 分别乘三组训练参数得到 Q、K、V,再按注意力头拆成[B,H,T,Dh]。Q 与 K 的转置相乘,得到[B,H,T,T]的位置匹配分数;分数除以√Dh,加入 Mask,并在 Key 位置上经过 Softmax,得到注意力权重。权重乘 V 后得到各头的上下文向量,再把所有头拼接并通过 WO 投影回[B,T,D]。随后经过残差、归一化和 FFN,主形状仍保持[B,T,D],因此可以连续堆叠多层。最后一层把 D 维表示映射成[B,T,N_vocab]的词表 logits,生成时使用最后一个有效位置预测下一个 Token。
十六、自测题
- Token ID 的
[B,T]怎样变成[B,T,D]? - 为什么同一个 X 可以生成 Q、K、V,但三者通常不相等?
- WQ 与 Q 的区别是什么?哪些是模型参数,哪些会随输入变化?
- 为什么要计算
Q×Kᵀ,而不是直接计算Q×K? [T,T]注意力矩阵的行和列分别表示什么?- 为什么缩放因子是
√Dh? - Mask 为什么必须在 Softmax 之前加入?
A [T,T]乘V [T,Dh]为什么得到[T,Dh]?- 多头输出为什么必须拼接并投影回 D 维?
- LayerNorm 与 RMSNorm 的主要差异是什么?
- 注意力与 FFN 分别主要负责什么?
- Pre-Norm 和 Post-Norm 的归一化位置有什么不同?
- 注意力 Softmax 与词表 Softmax 分别对什么对象归一化?
- Prefill 和 Decode 阶段的注意力形状有什么不同?
- KV Cache 为什么只保存历史 K、V,而通常不保存旧 Q?
- KV Cache 解决了什么问题,又没有解决什么问题?
