第 10 课:多头注意力 —— 让模型「多角度」看世界
第 10 课:多头注意力 —— 让模型"多角度"看世界
代码位置:src/attention.rs(
MultiHeadAttention结构 +forward) 底层算子:src/tensor.rs(reshape/permute/matmul/softmax_last_dim)
1. 本课要搞懂的问题
- 为什么"一个头"不够?多个头到底多学到了什么?
- 拆头为什么是 reshape → permute → reshape 三步,不能一步到位?
- 合并头之后为什么还要一个输出投影
c_proj? MultiHeadAttention::forward的 8 步,每一步形状怎么变?
2. 动机:一个头只能学一种"相关性"
第 9 课的注意力,一个查询位置对所有键算一组权重——它只能表达一种"谁和谁相关"的模式。
但文本里的相关性是多种多样的:
- 相邻词之间的语法依存("红 的 苹果")
- 指代关系("它" 指谁?)
- 全局语义(某个主题词反复出现)
一个头顾不过来。多头的思路:把 D 维空间切成 H 份,每份独立算注意力,让不同头学到不同模式,最后再合并。
| 头 | 可能学到的模式(示意,非真实数据) |
|---|---|
| 头 0 | 紧邻的前一个 token(位置关系) |
| 头 1 | 句法角色(主语-动词-宾语) |
| 头 2 | 指代消解(代词 ↔ 名词) |
| 头 3 | 全局主题词 |
数学形式:
MultiHead(Q, K, V) = Concat(head_1, ..., head_H) · W_O
head_i = Attention(Q·W_i^Q, K·W_i^K, V·W_i^V)
3. 结构一览
struct MultiHeadAttention {
c_q: Linear, // [D, D]
c_k: Linear, // [D, D]
c_v: Linear, // [D, D]
c_proj: Linear, // [D, D] 输出投影
n_head: usize,
}
数据流:
x → c_q / c_k / c_v 投影 → 拆头 → 每头各自算注意力 → 合并 → c_proj → 输出
n_head 来自配置 GPTConfig::tiny:n_embd = 64,n_head = 4 → head_dim = 64 / 4 = 16。
4. 拆头:reshape + permute + reshape
// 3. 拆头:[B, T, D] -> [B*H, T, head_dim]
// (先 reshape 出 H 维,再 permute 把 H 提到第 2 维)
let q = q
.reshape(vec![b, t, self.n_head, head_dim]) // [B,T,D] -> [B,T,H,head_dim]
.permute(&[0, 2, 1, 3]) // -> [B,H,T,head_dim]
.reshape(vec![b * self.n_head, t, head_dim]); // -> [B*H,T,head_dim]
为什么不能直接 reshape 成 [B*H, T, head_dim]?
关键在内存布局。以 D=8、H=2、head_dim=4、单个 token 为例,x 最后一维的数据排列是:
[ 头0的4维 | 头1的4维 ] ← 头 0 在前、头 1 在后(行优先存储)
reshape([B, T, H, head_dim]):只是重新解释数据,头 0 / 头 1 依然按块排列 ✓- 但如果跳过 permute 直接 reshape 成 [B*H, T, head_dim],B 维会按
b0-t0-h0, b0-t0-h1, b0-t1-h0, ...的顺序切片——同一个 token 的多个头被混进 T 维,完全错乱 ✗
所以必须先 permute(&[0, 2, 1, 3]) 把"头"维提到第 2 维(H 排在 T 前面),让每个 (b, h) 的数据在内存里连成一块,再 reshape 成 [B*H, T, head_dim],才能保证每个"批"恰好是一个头。
permute和reshape都不搬动数据,只改变"维度解释";三步合起来的效果 = "把 (b, h) 提成独立的批"。 这正是第 1 课讲的行优先存储的实战应用。
K、V 的拆头
K、V 的拆法完全一样(注意这里的 T 变成了 t_total——带 KV cache 时更长,第 18 课讲解):
let k = k
.reshape(vec![b, t_total, self.n_head, head_dim])
.permute(&[0, 2, 1, 3])
.reshape(vec![b * self.n_head, t_total, head_dim]);
// v 同理
为什么"先整体投影再拆头"是等价的?
c_q 是 [D, D] 线性层。对头 h 来说,它取输出的第 h 块(head_dim 维),这正好是输入 x 经过 W_Q 第 h 块列的作用——数学上等价于每个头一个独立的投影 W_h^Q(GPT-2 的实现也是这个思路,只是把投影合并写在一起)。
5. 并行注意力:一个批量 matmul 算完所有头
拆头后 [B*H, T, head_dim] 恰好是 matmul 支持的 3D 批量形状,第 9 课的公式一行不多、一行不少:
// 4. 注意力分数:scores = Q·Kᵀ / √d_k
let scale = 1.0 / (head_dim as f32).sqrt();
let kt = k.permute(&[0, 2, 1]); // [B*H, head_dim, T_total]
let scores = q.matmul(&kt).mul_scalar(scale); // [B*H, T, T_total]
// 5. 因果掩码:把"未来位置"变成 -inf,softmax 后概率为 0
let scores = scores.add(mask);
// 6. softmax 得到注意力权重,加权求和
let attn = scores.softmax_last_dim(); // [B*H, T, T_total]
let out = attn.matmul(&v); // [B*H, T, head_dim]
H 个头互不干扰,只是被"压扁"进了批维——这是拆头最大的好处:一次批量运算同时算完 H 个独立的注意力。
6. 合并头:拆头的逆操作
// 7. 合并头回 [B, T, D]
let out = out
.reshape(vec![b, self.n_head, t, head_dim]) // [B*H,T,head_dim] -> [B,H,T,head_dim]
.permute(&[0, 2, 1, 3]) // -> [B,T,H,head_dim]
.reshape(vec![b, t, d]); // -> [B,T,D](H×head_dim 拼回 D)
注意合并的顺序与拆头镜像对称:
| 拆头 | 合并 | |
|---|---|---|
| reshape | [B,T,D] → [B,T,H,head_dim] | [B*H,T,head_dim] → [B,H,T,head_dim] |
| permute | [0,2,1,3](H 提前) | [0,2,1,3](H 归位) |
| reshape | [B*H,T,head_dim](压平 B·H) | [B,T,D](拼回 D) |
拼回后,每个 (b, t) 的 D 维 = [头0的16维, 头1的16维, 头2的16维, 头3的16维],与拆头前的布局完全一致——因为 matmul 的批维上每个 (b, h) 都是独立的,头之间从不串扰。
7. 输出投影 c_proj:让模型决定"怎么混"
// 8. 输出投影
self.c_proj.forward(&out)
- 合并头只是"物理拼接",各头之间还没有交互
c_proj是一个 [D, D] 线性层:学习如何把 H 个头的信息融合(线性混合、重组)- 至此一个完整的多头注意力子层结束,输出形状 [B, T, D] 与输入一致,方便后面接残差连接(第 11 课)
数学上:
Concat(head_1..head_H) · W_O,W_O 就是c_proj的权重。
8. forward 全流程:8 步形状对照表
以 GPTConfig::tiny(D=64、H=4、head_dim=16)为例,B / T 视输入而定(训练演示时 B=8、T=block_size=32):
| 步骤 | 代码 | 形状 |
|---|---|---|
| 输入 x | x |
[B, T, 64] |
| ① Q/K/V 投影 | c_q/c_k/c_v.forward(x).reshape([b,t,d]) |
[B, T, 64] × 3 |
| ② KV cache(可选) | cache.append(...) |
K/V 的 T 变成 T_total(第 18 课) |
| ③ 拆头 | reshape → permute → reshape |
[B*4, T, 16] |
| ④ scores | q.matmul(&kt).mul_scalar(scale) |
[B*4, T, T_total] |
| ⑤ 掩码 | scores.add(mask) |
[B*4, T, T_total] |
| ⑥ softmax + 加权求和 | softmax_last_dim() → matmul(&v) |
[B*4, T, 16] |
| ⑦ 合并头 | reshape → permute → reshape |
[B, T, 64] |
| ⑧ 输出投影 | c_proj.forward(&out) |
[B, T, 64] |
完整的 forward 骨架(略去 KV cache 分支,逻辑与源码一致):
fn forward(&self, x: &Tensor, mask: &Tensor, kv_cache: Option<&mut KVCache>) -> Tensor {
let (b, t, d) = (x.shape()[0], x.shape()[1], x.shape()[2]);
let head_dim = d / self.n_head;
// ① 投影得到 Q、K、V
let q = self.c_q.forward(x).reshape(vec![b, t, d]);
let k = self.c_k.forward(x).reshape(vec![b, t, d]);
let v = self.c_v.forward(x).reshape(vec![b, t, d]);
// ③ 拆头
let q = q.reshape(vec![b, t, self.n_head, head_dim])
.permute(&[0, 2, 1, 3])
.reshape(vec![b * self.n_head, t, head_dim]);
let k = k.reshape(vec![b, t_total, self.n_head, head_dim])
.permute(&[0, 2, 1, 3])
.reshape(vec![b * self.n_head, t_total, head_dim]);
let v = v.reshape(vec![b, t_total, self.n_head, head_dim])
.permute(&[0, 2, 1, 3])
.reshape(vec![b * self.n_head, t_total, head_dim]);
// ④⑤⑥ 注意力(第 9 课)
let scale = 1.0 / (head_dim as f32).sqrt();
let kt = k.permute(&[0, 2, 1]);
let scores = q.matmul(&kt).mul_scalar(scale).add(mask);
let attn = scores.softmax_last_dim();
let out = attn.matmul(&v);
// ⑦ 合并头,⑧ 输出投影
let out = out.reshape(vec![b, self.n_head, t, head_dim])
.permute(&[0, 2, 1, 3])
.reshape(vec![b, t, d]);
self.c_proj.forward(&out)
}
9. 参数:训练时谁在被更新?
Module for MultiHeadAttention 负责收集全部可学习参数:
impl Module for MultiHeadAttention {
fn parameters(&self) -> Vec<Tensor> {
let mut ps = self.c_q.parameters(); // weight [64,64] + bias [64]
ps.extend(self.c_k.parameters());
ps.extend(self.c_v.parameters());
ps.extend(self.c_proj.parameters()); // 4 个 Linear = 8 个参数张量
ps
}
}
前向算出的误差通过反向传播(第 2 课)一路传回这 8 个张量,4 个投影层的权重随训练不断调整——"学什么相关性、怎么融合"全都由数据决定。
forward第 ② 步的 KV cache 只是推理时缓存历史 K/V 的加速手段,不改变计算结果,第 18 课专门讲解。
10. 运行与测试
cargo test # tensor 的 test_permute / test_matmul_3d 验证拆头、批量 matmul 所需算子
cargo run # 演示 3 训练小 GPT:4 头注意力参与每一轮前向/反向
11. 动手练习
- 把
GPTConfig::tiny的n_head从 4 改成 8(head_dim = 8)重新训练,观察 loss 是否变化(代码有assert_eq!(head_dim * self.n_head, d)保证必须能整除)。 - 手动构造一个 [2, 3, 8] 的 Tensor,按
reshape([2,3,2,4]) → permute([0,2,1,3]) → reshape([4,3,4])手写一遍,画出每个元素的归属,验证第 4 节的内存布局分析。 - 修改代码:去掉第 ⑧ 步
c_proj,直接返回合并后的 out,重新训练对比 loss(思考:表达能力损失在哪里)。 - 阅读
forward中kv_cache分支的cache.append,解释为什么推理时它能省掉"历史 token 的重复计算"(第 18 课预告)。
12. 本课总结
- 多头 = 把 D 切成 H 份,H 个注意力并行,学不同的相关性模式
- 拆头三步曲:reshape(切块)→ permute(H 提前)→ reshape(压平 B·H),顺序不能乱
- 合并是拆头的镜像;
c_proj负责融合各头信息 - 一次批量 matmul 同时算所有头,高效且简洁
- 下一步(第 11 课):位置编码 + LayerNorm + 残差连接,拼出完整的 Transformer Block!