深度学习

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

先不用急着理解每个符号。下面会从输入开始逐步解释。


二、统一符号

符号含义
BBatch Size,一次处理的样本数
TSequence Length,一条输入中的 Token 数
DModel Dimension,每个 Token 的主向量长度
HAttention Heads,注意力头数量
DhHead Dimension,每个注意力头的向量长度,标准 MHA 中通常 Dh=D/H
DffFFN 的中间层维度
N_vocab词表中的 Token 总数
NTransformer Block 的层数
X进入注意力模块的输入
QQuery,查询向量
KKey,匹配向量
VValue,内容向量
AAttention 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]

因为中间维度 23 不相等。

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]

BH 是批量维度,最后两个维度才是每次矩阵乘法真正使用的行和列。


五、从输入 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 被关注,它实际提供什么内容。

可以把一次注意力理解成:

  1. 用当前 Token 的 Q 去匹配所有 Token 的 K;
  2. 根据匹配结果得到注意力权重;
  3. 用这些权重组合所有 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

分数矩阵的阅读方式是:

当前 QueryKey 1Key 2Key 3
Token 1112
Token 2011
Token 3123
  • 每一行对应一个当前 Token 的 Query;
  • 每一列对应一个被匹配 Token 的 Key;
  • i 行第 j 列表示 Token i 的 Query 与 Token j 的 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=4Dh=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]
权重乘 VA×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。

十六、自测题

  1. Token ID 的 [B,T] 怎样变成 [B,T,D]
  2. 为什么同一个 X 可以生成 Q、K、V,但三者通常不相等?
  3. WQ 与 Q 的区别是什么?哪些是模型参数,哪些会随输入变化?
  4. 为什么要计算 Q×Kᵀ,而不是直接计算 Q×K
  5. [T,T] 注意力矩阵的行和列分别表示什么?
  6. 为什么缩放因子是 √Dh
  7. Mask 为什么必须在 Softmax 之前加入?
  8. A [T,T]V [T,Dh] 为什么得到 [T,Dh]
  9. 多头输出为什么必须拼接并投影回 D 维?
  10. LayerNorm 与 RMSNorm 的主要差异是什么?
  11. 注意力与 FFN 分别主要负责什么?
  12. Pre-Norm 和 Post-Norm 的归一化位置有什么不同?
  13. 注意力 Softmax 与词表 Softmax 分别对什么对象归一化?
  14. Prefill 和 Decode 阶段的注意力形状有什么不同?
  15. KV Cache 为什么只保存历史 K、V,而通常不保存旧 Q?
  16. KV Cache 解决了什么问题,又没有解决什么问题?

主要参考