Lesson 0003 · Phase 1 地基 · 约 45 分钟

Attention:从零手写

读完本课你将能:默写并亲手实现 scaled dot-product attention,解释 √d 和 causal mask 为什么必须存在

第一课的解剖图上,注意力层被标为"创新聚集地①"——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 的结果,进入饱和区:梯度趋近于零,训练不稳定。

回收第一课 这正是 QK-Norm 存在的理由:GLM-4.5 在"除 √d"之外又加了一道保险——直接对 Q、K 向量做 RMSNorm,把 logits 的范围摁得更死。模型一大,softmax 饱和问题就会回来,你已经在 config.json 里见过它的开关(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 部分亲手见到它。)

回收动画面板 C 六幕动画里灰色上三角格子写着 0——那是 softmax 之后落定的最终权重;而 -inf 是 softmax 之前在打分阶段塞进去的。两个阶段别混:先打分(-∞ 混在分数里)→ 再归一(0 分配落定)。这正是面板矩阵标题写"注意力权重 = softmax(QKT/√d)"的原因。

意外之喜:一个 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 训练快得多的结构根源——不是工程优化,是数学结构白送的。

边界:这个并行只属于训练 生成(推理)时,第 t 个 token 还没被生成出来,yt 依赖它——所以生成永远逐个来(第一课"逐 token 生成"那一格)。"训练时并行"的限定词就是在画这条边界:训练时整句话都是已知的教材,才可以一把算完。

四、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) × 38 套矩阵各自独立投影,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 个视角融成一个统一表达,交给下一层
输入 x —— 每个位置 512 维 (L, 512)
× 8 并行,各配一套独立 W_Q / W_K / W_V
头 1 → attn → 64 维 头 2 → attn → 64 维 头 8 → attn → 64 维

每头内部完整跑一遍第三节的流程(打分 → causal mask → softmax → 聚合),但只用自己那 64 维、画自己的 (L, L) 打分表

拼接 concat:8 × 64 = 512 (L, 512)
· W_O 混合 8 个视角 → 交给下一层 (L, 512)
回收动画面板 C 面板 C 画的单头 = 上图 8 个头里的任意 1 个(面板顶部"图为单头"那行说的就是它)。Multi-head 没有任何新机制——同一个公式跑 8 遍,只是 8 套 W 各学各的、8 张打分表各画各的。贵在视角多,不在机制新。
进阶 · 读完第五节脚本再回来看 跑完脚本你会发现 Multi-head 的实现核心就三行: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。
回收第一课伏笔:读者人数 × 笔记厚度,两个独立旋钮 "每位读者 64 维(512 ÷ 8)"只是惯例,不是定律。你亲手验证过的 GLM-4.5:请了 96 位读者(n_heads = 96)、每位笔记 128 维(head_dim = 128)——96 × 128 = 12288 ≠ 5120,注意力内维可以宽于隐藏维。读 config 的实用结论:n_heads(读者人数)和 head_dim(笔记厚度)是两个独立超参,别用 hidden ÷ 头数去脑补。这把钥匙以后看 DeepSeek / Kimi 的 config 时还会用到。

五、动手:可自判卷的实现

conda activate py310_qwenpaw
cd <课程工作区目录>
python scripts\attention_from_scratch.py

脚本干三件事(无需你写代码,但请逐行读懂 A 部分,注释里标了每步的 shape 流动):

  1. A · 从零实现:单头 scaled_dot_product_attention + 多头封装 MultiHeadAttention
  2. B · 对拍判卷:与朴素循环实现、torch 官方 F.scaled_dot_product_attention 双向对拍,误差已实测在 1e-16(机器精度)级,全绿 = 数学上完全等价
  3. C · 热力图:中文句「小猫追球因为它想玩」的 attention 热力图,输出到 scripts\attention_heatmap.png——左图用玩具语义嵌入(结构清晰),右图随机初始化(均匀混沌)对照
观察任务(带结果回来)
  1. 热力图左图:「它」这一行权重最大的是哪几个 token?接近 0 的是哪些?
  2. 「玩」这一行在看谁?和语义直觉一致吗?
  3. 左图第一行(小猫)为什么只有一个格子是 1.00?
  4. 右图和左图对比,说明"结构"来自哪里?

六、检索练习

七、本周深读任务(一手来源)

本周必做
  1. 动手(20 分钟):跑通脚本,读懂 A 部分每一行的 shape 变化,完成第五节的 4 个观察任务。
  2. 精读(30 分钟)Attention Is All You Need §3.2(Scaled Dot-Product Attention 与 Multi-Head Attention 两小节,原文仅 2 页)——本网络若访问不了 arXiv,用镜像或搜标题。读时对照本课公式,你会发现自己已经能"预告"论文的每一段。
  3. 可选加餐:Karpathy "Let's build GPT"——从零写完一个 GPT 全过程,attention 部分与本课互为印证。

读不懂 shape 流动、对拍出现 FAIL、热力图有意外发现——直接问我