Lesson 0003 · Phase 1 地基 · 约 45 分钟
Attention:从零手写
第一课的解剖图上,注意力层被标为"创新聚集地①"——MLA、GQA、QK-Norm 全都发生在这里。想看懂这些创新的"军备竞赛",得先走进赛场:本课从一行公式开始,到一段可验证的 PyTorch 实现结束。
一、为什么需要 attention:三条路线的动画对决
语言模型的本质矛盾:每个 token 需要汇聚"谁该看谁"的信息,但"谁该看谁"取决于内容本身。解决它有三条路线——先看动画,每条路线的"死穴"都藏在动画里。
路线一 · RNN:有内容感知,但路径是"挤牙膏"
要点:位置 i 唯一能看到的是 hi-1——一个固定尺寸向量的有损压缩。任意长的历史压不进固定大小,越早的 token 被稀释得越狠;训练时梯度要沿这条链走 L 步(消失/爆炸);时间步之间硬性串行,GPU 并行能力用不上。
路线二 · 全连接:路径直接,但连线是"焊死的"
要点:信息一步直达(路径 O(1)),但连线强度 wij 是训练学到的固定参数——对所有句子一视同仁,表达不了"这句话 3 该看 7、那句话不该"。且 wij 绑定坐标而非内容,换个长度全部失效。点"换一句话"看死穴。
路线三 · Attention:连线强度从"存储"变成"实时计算"
要点:每对 (i, j) 的连线强度 score = q·k 在前向传播时现算——WQ/WK 学的是"怎么算匹配",每句话的实际连线是当场重排的。直接路径(一步直达)+ 内容感知 + 全对并行(QKT 一次矩阵乘)三个性质同时拿到。
| 性质 | RNN | 全连接 | Attention |
|---|---|---|---|
| 路径长度(位置 1 → 100) | O(L),99 步压缩 | O(1) | O(1),一步直达 |
| 连线取决于内容? | ✓ 但要挤过瓶颈 | ✗ 焊死的 | ✓ 实时计算 |
| 训练可并行? | ✗ 串行链 | ✓ | ✓(一次矩阵乘) |
| 压缩瓶颈 | 有(固定 h) | 无 | 无(V 原样加权平均) |
代价:所有 (i,j) 对都算 = O(L²) 分数——后续长上下文创新(MLA 压 KV Cache、MoBA/NSA 稀疏化)优化的全是这张 L×L 矩阵。比喻:RNN = 传话游戏;全连接 = 焊死的电话交换机;attention = 每人举标签开麦、按匹配度实时分配的通信网。
所以第一课说注意力层是 "token 之间通信"——现在补全:是按内容动态寻址的通信。
二、一行公式,拆成四步流水
Attention(Q, K, V) = softmax( QKT / √dk ) · V
| 步骤 | 运算 | 直觉 |
|---|---|---|
| ① 相似度 | QKT → (L, L) 分数矩阵 | 每个 query 和每个 key 算内积——"我在找的"和"我贴的标签"有多匹配 |
| ② 缩放 | ÷ √dk | 防止分数过大(下一步会要命) |
| ③ 归一化 | softmax(·) 按行 | 分数 → 和为 1 的注意力分配 |
| ④ 聚合 | · V | 按注意力权重加权平均所有位置的 value——通信完成 |
为什么除 √dk?(本课最重要的"为什么")
假设 Q、K 的分量独立、零均值、方差为 1,则点积的方差等于 dk——维度越高,点积数值天然越大。大数值进 softmax 会发生什么?softmax 对悬殊的输入会输出接近 one-hot 的结果,进入饱和区:梯度趋近于零,训练不稳定。
use_qk_norm)。
三、Causal mask:不许偷看答案
语言模型是自回归的:训练时位置 i 的目标是预测第 i+1 个 token。如果 attention 让位置 i 看到了后面的 token,等于考试看答案——训练完全失效。
修法极其简单:在 softmax之前,把分数矩阵的上三角(j > i 的格子)置为 -inf,softmax 后这些位置的权重精确为 0。
为什么必须置 -inf,而不是"算完再删"?(这里有个反直觉的坑)
直觉方案是"把非法格子的分数设为 0,softmax 之后就当它不存在"。这个方案是错的,根源在 softmax 的机制:每个权重 = exp(分数) / Σ exp(所有分数)——注意分母里装着所有位置。而 exp(0) = 1:在 softmax 眼里,"0 分"不是一个被删除的格子,而是一个中等热度的普通分数。
用数字做一遍实验。假设某行允许的分数是 [2, 1],另有 1 个必须屏蔽的未来位置:
| 屏蔽方案 | softmax 输入 | 归一化结果 | 后果 |
|---|---|---|---|
| 设 0(直觉方案) | [2, 1, 0] |
(0.665, 0.245, 0.090) | 非法位置拿到 9% 的注意力——题漏了 9%;且它的 exp(0)=1 躺在分母里,把两个合法权重从 (0.731, 0.269) 压小——这就是"污染" |
| 置 -inf(正确方案) | [2, 1, -inf] |
(0.731, 0.269, 0) | 非法位置权重精确为 0,分母贡献也是 0——和"这个位置根本不存在"完全等价,合法权重分毫不动 |
所以 -inf 的意义是让"看不见"在数学上严格成立。(工程上写的是 -1e9 这类大负数——浮点数里 exp(-1e9) 就是 0,效果等价且避免真·无穷的运算问题。你跑脚本时会在 A 部分亲手见到它。)
意外之喜:一个 mask,换来训练全程并行
先回想 RNN 为什么快不起来(面板 A 的浑浊 h 链):它有状态链——h7 = f(h6),h6 = f(h5)……要算第 7 个位置的输出,必须先把前 6 步依次算完。长度 L 的序列 = L 个串行步,GPU 的并行能力在时间维度上完全用不上。
attention 的输出没有这条链。每个输出 yi 只依赖输入 x 和 mask,不依赖任何其他输出 yj:
y₁ = 用 x₁ 前缀算 attention ┐
y₂ = 用 x₁..x₂ 前缀算 attention │ L 行互不依赖
y₃ = 用 x₁..x₃ 前缀算 attention │ → 一个矩阵乘法同时算完
⋮ │
y_L = 用 x₁..x_L 前缀算 attention ┘
mask 在这里扮演的角色是让"并行算所有行"变得合法:第 i 行把 j > i 全部置 -inf,于是第 i 行自动只看自己的前缀——第 1 行看 1 个词,第 7 行看 7 个词,谁也没偷看未来。于是 L 行各自独立、互不等待,一次前向传播就把 L 个位置的输出全部拿到——等价于同时完成"用前 1 个词预测第 2 个词""用前 2 个词预测第 3 个词"……全部 L 个训练目标。对比 RNN 的 L 个串行步,这就是 Transformer 训练快得多的结构根源——不是工程优化,是数学结构白送的。
四、Multi-Head:同时请 8 位读者
第三节的一套 Q/K/V,整层注意力只有一张"谁看谁"的表——一种关注方式。但一句话里的关系是多样的:"它"指代谁(语义)、哪个词修饰哪个词(语法)、相邻词怎么搭配(位置)……一张表装不下这么多任务。
Multi-head 的做法:同时请 8 位读者读同一句话——每位只负责一种关系,各自画自己的"谁看谁"表,各自按表收集信息、写一份自己的阅读笔记,最后合并。下面把那句绕口的原文逐句翻译。
"投影到 d_head = hidden / n_heads 维子空间"=每位读者只带职责需要的信息
单头时代每个位置携带 512 维(hidden = 512)。请 8 位读者,每位不必扛着全部 512 维干活,而是只提取与自己职责相关的那个侧面:只管"指代关系"的读者,把句中与指代相关的线索提出来就够用了,64 维(d_head = 512 ÷ 8)装得下。这个"从 512 维全息信息里取职责相关那一面"的动作叫投影;得到的 64 维小世界叫子空间。
8 位各拿 64 维,合起来仍是 8 × 64 = 512——总算力不变,视角从 1 种变 8 种。这是 Multi-head 最划算的地方。
"各跑各的 attention"=互不商量,各画各的表
每位读者用自己学出来的阅读习惯(各自独立的一套 W_Q / W_K / W_V)把整句话重读一遍,画出自己的 (L, L) 表——8 头 = 8 张"谁看谁"表同时存在。没有谁规定 8 位要读出同样的东西:各自的阅读习惯被各自的梯度独立雕刻,自然走向不同分工(这个头盯指代、那个头盯相邻……)。然后各自按自己的表收集信息、写成 64 维阅读笔记——全程互不通信。
"最后拼接"=笔记装订 + 一次全体讨论
8 份 64 维笔记按位并排订成 512 维(拼接),再过一道 W_O 混合成每个位置最终的 512 维表达。8 位读者之间的交流只发生在 W_O 这一步——会诊时各说各的,主任医生最后综合。
一次前向传播的完整 shape 流动(hidden=512、8 头、d_head=64 为例)
| 步骤 | 运算 | 形状 | 要点 |
|---|---|---|---|
| 输入 | token 表达 x | (L, 512) | 每个位置一个 512 维向量 |
| 每头投影 | x · W_Q(h) / W_K(h) / W_V(h) | (L, 64) × 3 | 8 套矩阵各自独立投影,h = 1…8 |
| 每头打分 | q(h) · k(h)T / √64 | (L, L) × 8 | 每个头一张自己的"谁看谁"判断表——8 张表同时存在 |
| 每头聚合 | softmax(·) · v(h) | (L, 64) × 8 | 各头独立加权求和,互不通信 |
| 拼接 | concat 8 个头的输出 | (L, 512) | 每个位置上,8 个视角的发现并排站好 |
| 输出混合 | · W_O | (L, 512) | 把 8 个视角融成一个统一表达,交给下一层 |
每头内部完整跑一遍第三节的流程(打分 → causal mask → softmax → 聚合),但只用自己那 64 维、画自己的 (L, L) 打分表
q_all = x @ W_Q(一个 (512→512) 大矩阵一次算出全部 8 头的 q)→ q_all.view(L, 8, 64)(切开)→ 每头各自 attention。为什么"切输出"等价于"8 套独立投影"?矩阵乘法恒等式:(x @ W) 的第 h 块列 ≡ x @(W 的第 h 个 512×64 列块)——view 切出的每块,恰好等于"完整输入 × 该头专属小矩阵"。所以切的是输出(等价于切 W 的列),不是把输入 x 锯成 8 段——每头都完整看过 x。
五、动手:可自判卷的实现
conda activate py310_qwenpaw
cd <课程工作区目录>
python scripts\attention_from_scratch.py
脚本干三件事(无需你写代码,但请逐行读懂 A 部分,注释里标了每步的 shape 流动):
- A · 从零实现:单头
scaled_dot_product_attention+ 多头封装MultiHeadAttention - B · 对拍判卷:与朴素循环实现、torch 官方
F.scaled_dot_product_attention双向对拍,误差已实测在 1e-16(机器精度)级,全绿 = 数学上完全等价 - C · 热力图:中文句「小猫追球因为它想玩」的 attention 热力图,输出到
scripts\attention_heatmap.png——左图用玩具语义嵌入(结构清晰),右图随机初始化(均匀混沌)对照
- 热力图左图:「它」这一行权重最大的是哪几个 token?接近 0 的是哪些?
- 「玩」这一行在看谁?和语义直觉一致吗?
- 左图第一行(小猫)为什么只有一个格子是 1.00?
- 右图和左图对比,说明"结构"来自哪里?
六、检索练习
七、本周深读任务(一手来源)
- 动手(20 分钟):跑通脚本,读懂 A 部分每一行的 shape 变化,完成第五节的 4 个观察任务。
- 精读(30 分钟):Attention Is All You Need §3.2(Scaled Dot-Product Attention 与 Multi-Head Attention 两小节,原文仅 2 页)——本网络若访问不了 arXiv,用镜像或搜标题。读时对照本课公式,你会发现自己已经能"预告"论文的每一段。
- 可选加餐:Karpathy "Let's build GPT"——从零写完一个 GPT 全过程,attention 部分与本课互为印证。
读不懂 shape 流动、对拍出现 FAIL、热力图有意外发现——直接问我。