跳到主要内容

AttentionBlock 总装 — 一层注意力的完整流水线

这一章讲一件事: 第 04-06 章造的四个零件,按什么顺序、用什么姿势装成一层「会看上下文」的注意力块。 代码只有 30 行,但每一行的位置都有理由。读完你会拿到一张「张量形状漂流图」—— 一批文字进去,怎么变成一批「读过全文」的表示出来。

1. 先看现象:零件都对了,顺序错了也白搭

把第 06 章的 sdpa 函数直接套在文字表示上,会发生三件坏事: 数值大小没人管(训不动)、位置没人管(分不出词序)、每个词各配一套键值(显存爆炸)。 块(block)就是把「先稳定、再定位、再省着算」的顺序固定下来的容器。

先补本章唯一的新零件:残差连接(把块的输入原样加到块的输出上:输出 = 输入 + 变换(输入))1。 它管两件事:梯度可以沿着「原样加回去」那条短路直接流回浅层(第 02 章的深网训不动的老毛病,主要解法就是它); 模型只需要学「在原输入上什么」,不用学「从零生成全部」。

2. 主走查:一个张量走完整条流水线

这是本章的主走查:一批形状 [B, T, H] 的张量(B 句话并行、T 个词、H=256 维表示) 从头走到尾(注释里的含义标注全部对应该书代码)2:

x [B, T, 256]

├─ ① 存一份残差副本 residual = x
├─ ② RMSNorm(x) ← 预归一化:先把信号捋稳(第 04 章)
│ [B, T, 256]
├─ ③ qkv_proj(x) 一个线性层同时算出 Q|K|V 拼在一起
│ [B, T, 256+128+128] ← 4 查询头×64 + 2 KV头×64 ×2
├─ ④ 切三段、重排成 GQA 形状
│ Q [B,T,2,2,64] K/V [B,T,2,64]
│ ↑ 2 个 KV 头,2 个查询共用一套(第 06 章)
├─ ⑤ RoPE(q, k) ← 只转 Q 和 K,V 不转(第 05 章)
├─ ⑥ sdpa(q,k,v,sink,滑窗) ← 打分→掩码→加权(第 06 章)
│ [B, T, 4×64] → 拼回头维度
├─ ⑦ out_proj 把 4 个头的结果混回 256 维
│ [B, T, 256]

return residual + out ← 残差把输入原样加回

图说:出口形状 = 入口形状。这是块能堆几十层的结构性前提:
每一层都在「改写」同一批表示,而不是换一批。

三个「为什么在这个位置」,原书都给了理由:

归一化为什么在最先? 这是预归一化(pre-norm)的姿势:先捋稳再进零件, 深堆时梯度不会在入口就爆掉;LLaMA 一系全是这个顺序3。 块内不再放第二个归一化——稳信号的工作在下一个块开头还会做一次4

QKV 为什么用一个大投影而不是三个小的? 一个 [256 → 512] 的矩阵乘法在 GPU 上一次做完, 比三次小乘法少两遍调度开销;切三段是零拷贝的视图操作5。顺带一提,三个投影都不带偏置(偏置:额外加的一个常数项)—— 原书引的结论:偏置对表达力的贡献可忽略,纯增参数6

V 为什么不转? RoPE 只转 Q 和 K。打分发生在 Q·K,位置信息只影响「谁和谁相关」; V 是被掺走的内容本身,转了反而破坏语义7

3. 偶数层看近处,奇数层看全局

第 06 章说滑窗省算力但丢远信息。本书(和 Mistral)的答案是交替:

滑窗值 = 配置值,若 层号是偶数;否则 0(=全量注意力)

图说:偶数层只看最近 w 个词(便宜),奇数层看全文(保远信息)。
代码就一行:sliding_window = config.sliding_window if layer_idx % 2 == 0 else 0[^8]

原书给这笔账的意义:每两层里只有一层付全价, 整体注意力成本从「每层 O(T²)」压向「平均接近 O(T·w)」, 同时全局依赖仍有层可以走——这是「省」与「全」之间的量产折中8

4. 装配前的两道保险

构造函数里有两个整除检查,报错信息写得很直白9:

检查为什么必须
查询头数 % 键值头数 == 0GQA 要把查询头均分成组;4÷2=2 组,除不尽就切不匀
head_dim 是偶数RoPE 要把维度两两配对旋转;奇数个维度必有一个落单

这类检查写在构造时而不是训练崩溃时,是原书代码反复示范的习惯: 形状错误在装配现场报,比在几小时训练的深处报便宜得多。

5. 训练时与推理时,同一个块两种活法

原书专门对照了块在两个阶段的差异10:

阶段输入滑窗层sink_logits
训练整段序列一次进(并行算所有位置)每层按层号开/关当普通参数一起训练,从全 0 学起
推理一次一个新词同样交替学好的值固定,KV 缓存配合滑窗滚动了旧词,sink 保住上下文锚点

推理那行值得展开:生成时每个新词的 K/V 追加进缓存,窗口外的旧 K/V 可以驱逐(这正是滑窗省钱的地方); 而开头那几个被全局关注的词由 sink 兜住注意力,驱逐它们也不再崩11。 原书引的口径是:滑窗+sink 让流式生成撑到百万级 token 而质量不塌12

6. 作者的判断与证据

说法书里的证据我们的标注
预归一化是现代默认点名 LLaMA 的做法;并解释了它对深堆的意义3领域共识
偶/奇交替来自 Mistral原书明说 inspired by Mistral 的交替设计8有明确出处
融合 QKV 与去偏置是通用省法原书从 GPU 调度与参数量两个角度论证56工程共识
残差让深堆可行梯度沿恒等通路直通的机制描述1领域共识,与第 02 章的正则化一节呼应

判断(我们的,不是书里的): 这一章最值得带走的不是任何单个零件,而是**「形状不变」这个设计约束**—— 入口 [B,T,H]、出口 [B,T,H],残差和预归一化都在为它服务。正因为每一层不改形状, 层数才成了纯粹的「堆叠」问题,模型的「深度」才能变成一个配置数字(第 04 章)。 如果错,会错在: 如果某天主流架构改成了逐层变宽/变窄的设计(历史上确实有过), 「形状不变→深度只是堆叠」这条就不再普遍成立——但在 transformer 一族内它至今是事实。

7. 边界与局限

  • 本章的块不含前馈层:完整的「一层」还要再串一个 MLP 块——那是第 08 章的事,原书在章末预告了同样的顺序13;
  • 滑窗的窗口大小在本书配置里是固定数,原书没讨论逐层不同的窗口;
  • 推理时的 KV 缓存管理(驱逐、追加)在本书里是外部组件,块本身不管缓存——原书明说「Uses KV cache (external)」14;
  • 原书 42:493 一段说现代模型「通常 32-80 层」,与其 50 章说 GPT-4/PaLM「数百层」的口径不一致,取保守的 32-80 为准。

8. 可带走的

  1. 块 = 把「稳定→定位→省算→混合」的顺序固化;零件对、顺序错也白搭;
  2. 残差连接:输出 = 输入 + 变换(输入),梯度的高速公路,深堆的前提;
  3. 预归一化:RMSNorm 站在块的第一行;块内不放第二个归一化;
  4. QKV 一个大投影算完、不带偏置、零拷贝切片——三处都是省法;
  5. RoPE 只转 Q 和 K,V 不转:位置管「谁相关」,不管「内容本身」;
  6. 偶数层滑窗、奇数层全局:省一半以上的注意力钱,远信息仍有通路(Mistral 式交替);
  7. 构造时做整除检查:形状错误死在装配现场,不要死在训练深处;
  8. 推理时滑窗负责「缓存瘦身」、sink 负责「锚点保命」,两者配合才有长流式生成。

9. 原文地图

主题原书章原文位置
注意力三组件复习Fundamentals of Attentiontext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:21(搜「queries (representing the current focus)」)
自注意力的五条性质Self-Attention in Transformerstext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:51(搜「Permutation Equivariance」) · text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:55(搜「Parallel Processing」)
因果自注意力Self-Attention in Transformerstext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:71(搜「causal self-attention」)
sink 现象Limitations and Challengestext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:113(搜「Attention Sinks」)
多头各管一摊Multihead Attention in Transformerstext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:183(搜「Specialization of Heads」)
GQA/滑窗/sink 三定义What Is Grouped Query Attention, Sliding Window Attention, and Sink Tokens?text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:317(搜「Grouped Query Attention (GQA)」) · text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:329(搜「Sliding Window Attention (SWA)」) · text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:339(搜「Sink Tokens」)
构造函数与两道校验Integration of Attention Mechanism in Our Custom Large Language Modeltext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:361(搜「layer_idx」) · text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:375(搜「divisible」) · text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:45(搜「even」)
偶数层滑窗Integration of Attention Mechanism in Our Custom Large Language Modeltext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:383(搜「layer_idx % 2」)
融合 QKV 与无偏置Exhaustive Explanation of the AttentionBlock PyTorch Moduletext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:693(搜「Single matrix multiply」) · text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:591(搜「No bias term」)
sink_logits 全零初始化Exhaustive Explanation of the AttentionBlock PyTorch Moduletext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:633(搜「sink_logits」) · text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:635(搜「initialized to zeros」)
forward 全流程Integration of Attention Mechanism in Our Custom Large Language Modeltext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:439(搜「forward」) · text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:451(搜「self.norm(x)」) · text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:477(搜「self.rope(q, k)」) · text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:485(搜「residual + out」)
V 不转的理由Exhaustive Explanation of the AttentionBlock PyTorch Moduletext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:717(搜「Values (V) remain unrotated」)
GQA 4 倍缓存账Exhaustive Explanation of the AttentionBlock PyTorch Moduletext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:525(搜「4x smaller」)
交替设计的出处与账Exhaustive Explanation of the AttentionBlock PyTorch Moduletext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:549(搜「alternation of sparse」) · text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:839(搜「alternates with full attention」)
训练/推理对照Exhaustive Explanation of the AttentionBlock PyTorch Moduletext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:781(搜「Training:」) · text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:791(搜「Inference:」)
KV 缓存是外部组件Practical Implementation Notestext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:851(搜「KV caching」)
FlashAttention 定位Performance Optimizationstext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:815(搜「FlashAttention」)
章末预告 MLPSummarytext/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:861(搜「MLPBlock」)

Footnotes

  1. 出处:「Exhaustive Explanation of the AttentionBlock PyTorch Module」(Residual 段)(text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:661,搜「Residual connections」)。原文:残差连接让信息直通、防止深堆(如 80 层)中的梯度消失;forward 末行 return residual + out(:485)。 2

  2. 形状漂流图逐行对应该书代码(text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:451,搜「self.norm(x)」;:453,搜「qkv_proj」;:469(搜「reshape」)备选;:485,搜「residual + out」)。维度换算用第 04 章配置:4 查询头×64=256、2 KV 头×64=128。

  3. 出处:「Normalization Layer」(text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:557,搜「pre-norm」)。原文:RMSNorm 以 pre-norm 姿势放在注意力之前,通过保持输入方差一致来稳定深模型的梯度;并注明 LLaMA 采用。 2

  4. 出处:「Block Composition and Information Flow」(第 09 章文本)(text/56-fm-block-composition-and-information-flow.txt:5,搜「No terminal norm」)。原文:块内不放终端归一化,由后续块或最终层接手,pre-norm 链成端到端稳定。

  5. 出处:「QKV Projection Layer」(text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:693,搜「Single matrix multiply」)。原文:单个线性层算出拼接的 QKV,利用 GPU 的融合矩阵乘法,比分次投影少调度开销;切片是零拷贝视图(:693,搜「zero-copy」)。 2

  6. 出处:「Query-Key-Value Projections」(第 09 章文本)(text/51-fm-query-key-value-projections-and-multihead-decomp.txt:3,搜「bias-free」)。原文:刻意省掉偏置——偏置对表示能力的贡献可忽略,却给每个块平增参数。 2

  7. 出处:「RoPE Application」(text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:717,搜「Values (V) remain unrotated」)。原文:V 不旋转——位置影响相似度分数(Q·K),不影响取回的内容。

  8. 出处:text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:549(搜「alternation of sparse」)与 :805(搜「alternates with full attention」)。原文:稀疏(滑窗)与稠密(全量)交替、灵感来自 Mistral,兼顾局部建模与全局依赖,可扩到百万级 token。 2

  9. 出处:text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:375(搜「divisible」)与 :377(搜「even」)。原文两处 ValueError:查询头数必须被键值头数整除;RoPE 要求 head_dim 为偶数。

  10. 出处:「Training and Inference Behaviors」(text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:781,搜「Training:」;:791,搜「Inference:」)。原文:训练整段进、sink_logits 随反传学习;推理用外部 KV 缓存追加新词,SWA+sink 支撑流式。

  11. 出处:「Sink Logits」(text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:635,搜「initialized to zeros」)与 :811(搜「contextual anchors」)。原文:sink 给开头 token 加偏置使其成为上下文锚点;流式驱逐旧词时保住缓存稳定性。

  12. 出处:text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:793(搜「streaming」)。原文:「SWA + sink logits enable streaming (e.g., 4M+ tokens) without perplexity collapse」。

  13. 出处:「Summary」末段(text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:861,搜「MLPBlock」)。原文:第 8 章转向注意力之后的馈送部件 MLPBlock,并预告 MoE 与 SwiGLU。

  14. 出处:「Practical Implementation Notes」(text/42-ch07-7-attentionblock-with-rotary-embedding-gqa-slidi.txt:851,搜「KV caching」)。原文:推理用 KV 缓存(未在块内实现)存过去 K/V,逐步追加;sink 保缓存稳定、SWA 限制缓存增长。