第 10 课:多头注意力 —— 让模型「多角度」看世界

2026-09-11 干徒
RustLLM注意力

第 10 课:多头注意力 —— 让模型"多角度"看世界

代码位置:src/attention.rsMultiHeadAttention 结构 + forward) 底层算子:src/tensor.rsreshape / permute / matmul / softmax_last_dim

1. 本课要搞懂的问题

  1. 为什么"一个头"不够?多个头到底多学到了什么?
  2. 拆头为什么是 reshape → permute → reshape 三步,不能一步到位?
  3. 合并头之后为什么还要一个输出投影 c_proj
  4. 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],才能保证每个"批"恰好是一个头。

permutereshape 都不搬动数据,只改变"维度解释";三步合起来的效果 = "把 (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. 动手练习

  1. GPTConfig::tinyn_head 从 4 改成 8(head_dim = 8)重新训练,观察 loss 是否变化(代码有 assert_eq!(head_dim * self.n_head, d) 保证必须能整除)。
  2. 手动构造一个 [2, 3, 8] 的 Tensor,按 reshape([2,3,2,4]) → permute([0,2,1,3]) → reshape([4,3,4]) 手写一遍,画出每个元素的归属,验证第 4 节的内存布局分析。
  3. 修改代码:去掉第 ⑧ 步 c_proj,直接返回合并后的 out,重新训练对比 loss(思考:表达能力损失在哪里)。
  4. 阅读 forwardkv_cache 分支的 cache.append,解释为什么推理时它能省掉"历史 token 的重复计算"(第 18 课预告)。

12. 本课总结

  • 多头 = 把 D 切成 H 份,H 个注意力并行,学不同的相关性模式
  • 拆头三步曲:reshape(切块)→ permute(H 提前)→ reshape(压平 B·H),顺序不能乱
  • 合并是拆头的镜像;c_proj 负责融合各头信息
  • 一次批量 matmul 同时算所有头,高效且简洁
  • 下一步(第 11 课):位置编码 + LayerNorm + 残差连接,拼出完整的 Transformer Block!
Rust 大语言模型 学习指南共 22 章
1从零用 Rust 实现大语言模型 —— 学习计划2第 1 课:张量 Tensor —— 一切的基础3第 2 课:自动微分 Autograd —— 让模型学会「自我修正」4第 3 课:张量运算扩展 —— 广播、归约、softmax、批量矩阵乘法5第 4 课:模块化重构 —— 项目结构分层与 Rc<RefCell> 架构6第 5 课:线性层与激活函数 —— 神经网络的「积木」7第 6 课:损失函数与优化器 —— 让模型知道「错在哪、怎么改」8第 7 课:第一个 MLP —— 教会神经网络算 XOR9第 8 课:BPE 分词器 —— 让模型「读懂」文字10第 9 课:注意力机制 —— 让 token 互相「看」11第 10 课:多头注意力 —— 让模型「多角度」看世界本篇12第 11 课:位置编码与归一化 —— 让序列带上「位置感」13第 12 课:完整 GPT 模型 —— 把积木拼成能预测下一个词的模型14第 13 课:训练循环 —— 让模型真正开始学习15第 14 课:数据加载 —— 文本如何变成训练样本16第 15 课:推理与采样 —— 让模型「创造性」地生成17第 16 课:训练小 GPT —— 看 loss 从 1.46 一路降到 0.1618第 17 课:AdamW 优化器 —— 给梯度下降装上「惯性」和「自适应步长」19第 18 课:KV Cache —— 让逐 token 生成不再重复计算20第 19 课:RoPE 旋转位置编码 —— 把「相对位置」揉进注意力21第 20 课:学习率调度与收尾 —— warmup、cosine decay 与全项目总结22第 21 课:GPU 加速训练与推理