跳到主要内容

数据截至 (上游 commit eb980a5c9eea)

手搓一个 LLaMA2 — 今天的解码器换掉了哪四个零件

这一章讲三件事: 2017 年的原版和今天的模型到底差在哪;每一处替换是为了解决什么问题、 代价是什么;以及一个模型「吐字」的那几行代码里,温度(就是控制它敢不敢挑冷门词的那个旋钮)和 top-k 分别在干什么。 这是全书从「读懂」转向「动手」的第一章。 后面两章训练用的模型,就是这一章搭出来的这个。 不需要会写代码——我们不抄代码,只讲每个零件在干什么、为什么。

1. 先看现象:课本上的 Transformer,和今天的模型不是一个东西

你如果照着第 03 章那张装配图去看今天任何一个开源模型的代码,会发现对不上:

第 03 章讲的今天真正在跑的
层归一化(Layer Norm)RMSNorm
正余弦位置编码,加到词向量上旋转位置编码,把向量转个角度
每个注意力头各有一组 Q、K、V多个查询头共用一组键值头
前馈层 = 两层线性 + ReLU三层线性 + 门控

四处都换了。而且换的理由各不相同。

这一章就是把这四处替换讲透。 书里的载体是 LLaMA2——Meta 2023 年 7 月发布的开源模型1, 而这四处替换今天几乎是所有开源模型的标配。

2. 顶层全景:一台今天的解码器

一句话
└ 分词器 → 编号
└ 查表层 → 向量

┌────┴──── 一层解码器(重复 N 层)────────────────┐
│ ① RMSNorm │
│ ② 注意力(带因果掩码 + 分组查询 + 旋转位置编码) │
│ ③ 加回输入(残差) │
│ ④ RMSNorm │
│ ⑤ 门控前馈层 │
│ ⑥ 加回上一步(残差) │
└───────────────────────────────────────────────┘

└ 最后一次 RMSNorm
└ 线性层 → 词表大小的一串数 → 挑一个 → 接到句尾,整条流水线从头再跑

图说:骨架和第 03 章完全一样(归一化 → 子层 → 加回原输入,重复两遍)。
换掉的是每个格子里具体装的东西。

书里给的示范配置2:

设定意思
dim768隐藏层维度——每个标记用 768 个数表示
n_layers12上面那个方框重复 12 次
n_heads16注意力切成 16 个头
n_kv_heads8但键值只有 8 组——这就是分组查询,见第 5 节
vocab_size6144词表里有 6144 个标记
max_seq_len512一次最多吃 512 个标记

按这个配置搭出来的模型,参数量约 8259 万3第 07 章真正训练时用的是一个更大的版本,2.15 亿参数。

主走查:一句话穿过这四个换掉的零件

四个零件各讲各的,读者读完只会记住四句「换成了什么」。所以这一章挑一句话,穿过全部四个:

「我喜欢你」——第 03 章那句话,切成 3 个标记。 配置就用上面那张表:隐藏层维度 768,16 个查询头、8 组键值头,所以每个头分到 768 ÷ 16 = 48 维。

后面第 3 到第 6 节,每一节末尾都有一小段**「我喜欢你」走到这个零件**, 写出它此刻的形状和至少一个具体的数。 全章走完是这四步:

查表 → 3 行 × 768 个数
① RMSNorm → 3 × 768(形状不变,每行被除以自己的均方根)
② 旋转位置编码 → 只作用在查询和键上,第 n 个位置转 n 份角度
③ 分组查询注意力 → 查询 16 组 × 48 维,键值只有 8 组 × 48 维
④ 门控前馈层 → 每行 768 → 2048 → 768(中间那个 2048 是算出来的,见第 6 节)

图说:进出都是 3 × 768,和第 03 章那台机器一模一样;换的全是格子里的东西。

先声明:下面这些小段里的数值,凡不是书里给的,全是为演示编的,不是真实数值。 上面那张配置表(768 / 16 / 8 / 512)来自书里。

3. 换掉零件一:RMSNorm——省掉一次求平均

这一节回答:第 03 章那个层归一化,为什么还能再简化。

先回忆层归一化在做什么

第 03 章讲过:减去平均值(叫「居中」),再除以标准差(叫「缩放」)。 两步各有各的作用。

RMSNorm 只做第二步

RMSNorm 的做法:不减平均值,直接除以「均方根」。

「均方根」是什么: 把一组数各自平方、取平均、再开方。 它衡量的是这组数整体有多大,但不关心它们的中心在哪。

书里对它的说明4:

  • 有一个可学习的缩放参数(和层归一化那两个参数里的一个对应);
  • 分母上加一个小常数,避免除以零;
  • 作用是「确保权重的规模不会变得过大或过小来稳定学习过程,这在层数很多的模型里特别有用」。

为什么可以省掉「减平均值」

书里没有回答这个问题,只是照着 LLaMA 的做法实现。我们补上。

补充(不在书里):RMSNorm 的原始出处是 2019 年 10 月的论文。 它的假设写得很直白:层归一化里的「重新居中不变性」是可有可无的, 真正起作用的是「重新缩放不变性」。 去掉居中之后,在不同模型上省下 7% 到 64% 的运行时间5

层归一化: (x - 平均值) / 标准差 × 缩放 + 平移 ← 两个参数,要算两个统计量
RMSNorm : x / 均方根 × 缩放 ← 一个参数,只算一个统计量

图说:少算一个统计量、少一次遍历、少一个参数。
在一个 12 层的模型里,这道手续每层要做两次,省下来的是实打实的时间。

判断(我们的,不是书里的): 这处替换值得记住,不是因为它多聪明, 而是因为它示范了这一行最常见的一类进步:把一个环节里「其实没在起作用的那一半」删掉。 论文的贡献不是发明了什么,是证明了少做一件事不会更差。 这类工作在论文里显得平淡,在工程上却直接省钱。 如果错,会错在: 如果在某些任务或某些深度上,居中这一步确实重要 (有研究认为它在特定初始化下影响训练早期的稳定性),那么「可有可无」就要加限定条件。

「我喜欢你」走到这一步(主走查 ①)

这一步对每一行自己的 768 个数做。 768 个手算不动,拿 4 个数演示同一套算术, 并且把第 03 章那套并排放着,差别就一目了然:

「我」那一行(取 4 个数):[ 3, -1, 2, 2 ]

第 03 章的层归一化:
平均值 = (3-1+2+2) ÷ 4 = 1.5
减掉平均值 → [ 1.5, -2.5, 0.5, 0.5 ]
标准差 = √((2.25+6.25+0.25+0.25) ÷ 4) = √2.25 = 1.5
再除以 1.5 → [ 1.00, -1.67, 0.33, 0.33 ]

这一章的 RMSNorm:
均方根 = √((9+1+4+4) ÷ 4) = √4.5 ≈ 2.12 ← 直接平方求平均再开方,不减任何东西
每个数除以 2.12 → [ 1.41, -0.47, 0.94, 0.94 ]

图说:上面那套要遍历两遍(先求平均值,再求标准差),下面这套只遍历一遍。
出来的两串数不一样,但「这一行整体多大」被拉到同一个尺度上了——
论文赌的就是「后面那件事才重要,前面那件事可有可无」。
这句话每过一层要做两次,12 层就是 24 次,每次都省一遍。
(这 4 个数是为演示编的,不是真实数值。)

4. 换掉零件二:旋转位置编码——从「加」改成「转」

这一节讲四处替换里最有巧思的一处。

先回忆问题

第 03 章讲过:注意力眼里没有语序,必须额外补进去。 原版的做法是给每个位置算一串数,加到词向量上。

这个做法有个隐患: 位置信息和词义信息被混在同一组数里, 模型要自己把它们分开。而且它编的是「绝对位置」——第 5 个字就是第 5 个字, 至于第 5 个字和第 8 个字隔了 3 位,得靠模型自己算出来。

旋转位置编码的做法:不加,转

核心想法一句话:

不给向量加东西,而是把它按位置转一个角度—— 第 1 个位置转 1 份角度,第 2 个位置转 2 份,以此类推。

为什么转角度就能表示相对位置:

两个向量的点积(第 02 章讲的「有多像」),只跟它们的夹角有关。
位置 m 的向量转了 m 份角度,位置 n 的向量转了 n 份角度
⟹ 它们的夹角差正好是 (m - n) 份
⟹ 算出来的注意力分数,天然只跟「隔多远」有关,跟「在哪」无关

图说:这就是它比「加一串数」高明的地方——
相对位置不是被模型学出来的,是被数学结构直接保证的。

书里的实现分三步

书里没有讲上面这个道理,直接给了三个函数6:

函数干什么
precompute_freqs_cis预先算好每个位置该转多少角度,以正弦和余弦两张表的形式存下来
reshape_for_broadcast调整形状,让这两张表能和实际的向量对上
apply_rotary_emb真正执行旋转——把查询和键都转一下

第一个函数的细节值得看一眼,因为它和第 03 章那个正余弦公式同源7:

① 生成一串频率:每一维用不同的「转速」,底数仍然是 10000
② 生成一串位置:0, 1, 2, …, 最大长度
③ 两者相乘 → 得到「第几个位置、第几维,该转多少角度」
④ 对这个角度分别取余弦和正弦,存成两张表

图说:和第 03 章的正余弦编码用的是同一套频率设计。
差别只在最后怎么用它——那边是加上去,这边是拿来转。

注意一个容易看漏的细节:旋转只作用在查询和键上,不作用在值上8

为什么: 位置信息是用来算「谁和谁相关」的,而查询和键正是算这件事的两个东西。 值是「内容」,内容本身不该被位置扭曲。

收益:长度外推

第 05 章提过:书里说大部分大模型用旋转位置编码,是因为它具有一定的长度外推能力—— 推理时能处理明显长于训练长度的文本9

为什么它能外推,而 BERT 那种查表法不能: 查表法只有 512 行,第 513 个位置没有对应的行; 旋转位置编码是一条公式,第 513 个位置的角度直接算出来就行。

补充(不在书里):旋转位置编码的原始出处是 2021 年 4 月的 RoFormer 论文, 摘要里点明的三个性质是:用旋转矩阵编码绝对位置、在自注意力里自带相对位置依赖、 以及随相对距离增大而衰减的依赖关系10

最后那条「随距离衰减」值得注意: 它意味着这套编码天然认为「离得远的关系弱」—— 这是一个被写进结构里的默认假设,不是模型自己从数据里学出来的。 换句话说,这条倾向从模型第一天开机就在,数据说不说得通都改不掉它。

「我喜欢你」走到这一步(主走查 ②)

「转一个角度」不是比喻,是真的在转。 48 维是一对一对转的,拿其中一对演示:

假设「每份角度」是 30 度。三个位置各转:我 0 度、喜欢 30 度、你 60 度。

「你」那个位置,某一对相邻维上的两个数是 (0.80, 0.60):
转 60 度之后 =
新的第一个数 = 0.80×cos60° - 0.60×sin60° = 0.40 - 0.52 = -0.12
新的第二个数 = 0.80×sin60° + 0.60×cos60° = 0.69 + 0.30 = 0.99
→ (0.80, 0.60) 变成 (-0.12, 0.99)

验一件事:转之前 0.80²+0.60² = 1.00,转之后 0.12²+0.99² ≈ 0.99。
└ 长度没变,只有方向变了 —— 这就是「转」和「加一串数」的根本区别。

「你」和「我」隔了 2 位:60° - 0° = 60°,正好是 2 份。
「你」和「喜欢」隔了 1 位:60° - 30° = 30°,正好是 1 份。
└ 这两句话就是「相对位置由数学结构直接保证」的全部含义:
谁在第几位不重要,重要的是两个角度一减,差出来的正好是「隔了几位」。

图说:注意这只对查询和键做,值不转 —— 值是内容,内容不该被转。
(30 度、(0.80, 0.60) 是为演示编的,不是真实数值;真实的每份角度逐维不同,
而且由第 03 章那套底数 10000 的频率算出来。)

5. 换掉零件三:分组查询注意力——省的是推理时的显存

这一节讲一个纯粹为了「跑得动」而做的妥协。

先看问题:生成的时候什么最贵

第 04 章说过,生成是一个标记一个标记来的:每吐一个字,整条流水线跑一遍。

但有一件事不必重跑:前面那些标记的键和值。 它们算过一次就不会变,所以实现里会把它们存下来重复使用—— 这套缓存有个名字叫 KV 缓存,中文叫键值缓存——键和值算过一次就存着,下次直接取。

问题就出在这里:

序列越长 → 缓存越大 → 每生成一个字,都要把整个缓存从显存里读一遍
└ 瓶颈不是「算得慢」,是「搬得慢」

图说:生成阶段真正卡住的是显存带宽,不是算力。
缓存里装的东西越少,搬得越快。

先把两笔账拆开:一笔是平方的,一笔是线性的

这里必须停一下,因为这两笔账最容易被混成一笔,而它们的形状完全不同。

第 02 章说过,注意力要算「每个位置对每个其他位置」的相似度, 所以算注意力分数这笔计算量,随序列长度的平方增长。

而这一节说的键值缓存是另一笔:每多一个位置,就多存一个位置的键和值—— 它随序列长度线性增长。

拿两个具体长度对照着看,差别一目了然:

序列长度算注意力分数的次数(每层每头)键值缓存里存了几个位置
512512×512 ≈ 26 万512
40964096×4096 ≈ 1678 万4096
长度涨 8 倍,这一栏涨了64 倍8 倍

两笔账,两种省法,别张冠李戴: 分组查询注意力省的是右边那一栏(缓存里少存几组键值), FlashAttention 改的是左边那一栏的搬运方式(计算量一次不少,只是少搬几趟)。 一个动的是「存多少」,一个动的是「怎么搬」,谁也替不了谁。

三种做法

做法有几组键值效果
多头注意力(原版)每个查询头配一组,16 个头就是 16 组质量最好,缓存最大
多查询注意力(MQA)所有头共用一组缓存最小,但质量掉得明显
分组查询注意力(GQA)分组共用,比如 16 个查询头共用 8 组折中

书里的实现里,16 个查询头配 8 组键值头11,也就是每两个查询头共用一组。

实现上的一个小函数

因为查询有 16 组、键值只有 8 组,算注意力之前要先把键值复制一份凑够数。 书里为此写了一个 repeat_kv 函数12:

键值张量: [批, 序列长, 8 组头, 每头维度]
└ 在「组」这一维后面插一个新维度
└ 把新维度扩展成「要重复几次」(这里是 2)
└ 再压平回去 → [批, 序列长, 16 组头, 每头维度]

图说:注意这只是「摆给注意力看」的形状,
真正存进 KV 缓存的仍然是 8 组——省显存的部分在那里。

「我喜欢你」走到这一步(主走查 ③)

把这三个字的账真的算一遍,省了多少就看得见了。 每个头 48 维,三个位置:

查询 Q: 3 位置 × 16 个头 × 48 维 = 2304 个数 ← 这一栏不省,查询照样是 16 组
键 K: 3 位置 × 8 组 × 48 维 = 1152 个数 ← 省在这里
值 V: 3 位置 × 8 组 × 48 维 = 1152 个数 ← 和上面一样

要存进 KV 缓存的 = K + V = 1152 + 1152 = 2304 个数
换成原版多头注意力(16 组键值)= 2304 + 2304 = 4608 个数
└ 省掉一半,而且一个数都不用少算 —— 算之前 repeat_kv 把 8 组各复制一份就补齐了

图说:三个字省 2304 个数看着不多;把序列拉到 512 个标记、12 层一起算,
省下的就是 2304 ÷ 3 × 512 × 12 ≈ 470 万个数。
而每生成一个字,这些数都要从显存里整个读一遍。
(48、16、8、12、512 全部来自第 2 节那张配置表,不是编的。)

书里在这里给了一个诚实的交代

LLaMA2 里只有 700 亿参数那一版用了分组查询注意力, 但书里在小模型上也用了它,理由是「可以提高模型的效率,并节省一些显存占用」13

补充(不在书里):分组查询注意力的原始出处是 2023 年 5 月的论文。 它的出发点正是上面那个问题:多查询注意力虽然大幅加快解码,但会损失质量; 分组查询用「多于一组、少于查询头数」的键值头做折中, 论文报告的结果是质量接近多头注意力,速度接近多查询注意力14

论文里还给了一个很实用的做法:已经训好的多头模型, 只要用原始预训练算力的 5% 再训一小段,就能转成分组查询的版本14

顺带:代码里那个 flash_attn 开关

书里的配置里有一个 flash_attn 开关,实现里也确实会去检测能不能用它15, 但正文一个字都没有解释它是什么。

补充(不在书里):它指的是 FlashAttention,2022 年 5 月的论文提出的方法。 它不改变注意力的计算结果,只改变数据在显存和片上缓存之间怎么搬—— 用分块的方式减少搬运次数,从而在完全不做近似的前提下明显加速16

和分组查询注意力的区别值得说清: 分组查询是改结构、换质量;FlashAttention 是改实现、不换质量。 一个是设计上的取舍,一个是纯粹的工程优化。

6. 换掉零件四:门控前馈层——多一条通道决定「放多少过去」

这一节讲四处替换里最不起眼、但最难解释的一处。

原版是什么样

第 03 章讲过:拉宽 → 过 ReLU → 缩回,两个线性层。

LLaMA2 是什么样

三个线性层17:

x ─┬─ 线性层 w1 → 拉宽 → 过 SiLU 激活 ─┐
│ × ← 两条通道逐位相乘
└─ 线性层 w3 → 拉宽 ────────────────┘

线性层 w2 → 缩回 → dropout

图说:多出来的那条 w3 通道叫「门」。
它逐位决定另一条通道的每个数「放多少过去」。

「门控」这个词的意思就是这个:一条通道当阀门,控制另一条通道的流量。

「SiLU」是什么: 又一个激活函数,形状介于 ReLU 和平滑曲线之间—— 负数不是硬压成 0,而是压成一个接近 0 的小负数。

前馈层维度的算法有点绕

先分清两个数,别混: 「隐藏层维度」是每个标记从头到尾被表示成的那串数有多长 (上面那张表里的 dim,768);而前馈层中间那一段会先把它临时拉宽再缩回来, 拉宽到多少,这里叫前馈层维度下面算的是后者。

书里的代码里有一段值得看18:

① 先按老规矩:前馈层维度 = 隐藏层维度 × 4
② 再乘 2/3 ← 因为多了一条通道,要把参数量拉回来
③ 再向上取整到 64 的倍数 ← 对齐显卡的计算单元

图说:第 ② 步是关键——三条通道各占 2/3 宽,总参数量和原来两条通道全宽差不多。
也就是说,门控是「白送」的,不是靠加参数换来的。

「我喜欢你」走到这一步(主走查 ④)

先把上面那三步按 768 算出来:

① 768 × 4 = 3072
② 3072 × 2/3 = 2048
③ 2048 已经是 64 的倍数(2048 ÷ 64 = 32),不用再往上取

所以这一层的前馈层维度是 2048。

「白送」这句话可以当场对账:

第 03 章两条通道、全宽 3072: 768 × 3072 × 2 = 471.9 万个参数
这一章三条通道、宽 2048: 768 × 2048 × 3 = 471.9 万个参数
└ 两个数一样。多出来的那条门通道,是靠「每条都窄一点」换来的,不是加出来的。

再看「我」那一行里的一个数,门是怎么关的:

「我」那一行 768 个数 → 两条通道各拉宽到 2048 个

盯住第 7 位:
w1 那条(要过 SiLU 的): -0.50 → 过 SiLU → -0.19
w3 那条(当门的): 2.00 → 原样
两条逐位相乘: -0.19 × 2.00 = -0.38 ← 这一位实际送出去的值

换一个数看门关小的样子:
w1 那条: -0.50 → -0.19 w3 那条: 0.05
相乘: -0.19 × 0.05 = -0.01 ← 门开到 0.05,这一位几乎被掐掉了

最后 2048 个数过 w2 缩回 768 个,这一行就算完了。

图说:门那条通道不做任何「思考」,它只决定另一条通道的每个数放多少过去。
同一个 -0.19,门开 2.00 就放大成 -0.38,门开 0.05 就掐成 -0.01。
(-0.50、2.00、0.05 是为演示编的,不是真实数值;768、4 倍、2/3、64 来自书里的代码。)

为什么有效:书里没说,而且原论文也没说

这一点必须挑明。

补充(不在书里):这类门控前馈层的出处是 2020 年 2 月的一篇短论文《GLU Variants Improve Transformer》, 作者 Noam Shazeer。它把前馈层里的激活函数换成各种门控变体,发现有几种确实更好19

判断(我们的,不是书里的):这是一个纯经验结论,没有机制解释。 原论文只报告了「换了之后质量更好」,没有解释为什么。 后来它被 LLaMA 采用、又被几乎所有开源模型跟进,主要依据是「大家都这么做,而且确实好一点」。 遇到这类零件,正确的态度是:用它,但别去给它编一个原理。 如果错,会错在: 如果后来出现了对门控机制的理论解释(比如它等价于某种「自动挑出该留哪几个数」的机制), 那这条「纯经验」的定性就要改。

7. 拼起来,以及吐字的那几行

一个解码器层

把上面四个零件按第 03 章的骨架装起来就完了20:

h = x + 注意力( RMSNorm(x) )
out = h + 门控前馈层( RMSNorm(h) )

和第 03 章那两行一模一样,只是括号里的东西全换了。 N 个这样的层叠起来,最后再做一次 RMSNorm,过一个线性层换算成一串和词表等长的数,就是完整模型。

吐字:生成函数在干什么

书里的生成函数只有二十来行,但它把第 04 章讲的「接龙」变成了可执行的动作21:

循环 max_new_tokens 次:
① 如果句子超过最大长度,把开头截掉
② 整条流水线跑一遍,只取最后一个位置的输出
③ 挑一个标记出来 ← 见下,这一步有讲究
④ 如果挑到了结束标记,停
⑤ 把它接到句尾,回到 ①

图说:第 ② 步是「整条流水线跑一遍」——这就是写长文必然慢的原因。

第 ③ 步:温度和 top-k

这两个参数你在任何一个模型接口里都见得到,这里把它们讲清。

模型算出来的是一串原始分数(每个标记一个),行话叫 logits。挑法有两种22:

温度怎么挑你会看到什么
等于 0直接挑分数最高的那个完全确定,同一个问题永远同一个答案;写长了容易重复
大于 0分数先除以温度,再转成可能性,然后按可能性随机抽有变化

「按可能性随机抽」这个动作,这一行叫采样——和「死挑分数最高的那个」相对。 下面说的「采样」都是指这一步。

为什么「除以温度」能控制随机程度:

原始分数: [ 5.0, 3.0, 1.0 ]
÷ 0.5 → [10.0, 6.0, 2.0 ] 差距被放大 ⟹ 转成可能性后,第一名几乎独占 ⟹ 更保守
÷ 2.0 → [ 2.5, 1.5, 0.5 ] 差距被压缩 ⟹ 后面的也有机会 ⟹ 更放飞

图说:温度小于 1 是「拉开差距」,大于 1 是「抹平差距」。
它本身不改变排名,只改变排名之间的悬殊程度。
(这三个原始分数是为演示编的,不是真实数值——书里这一段只有代码,没有示例数值。)

top-k 是另一道闸22:只保留分数最高的 k 个候选,其余全部排除,再在这 k 个里按可能性抽。 作用是堵住长尾——防止那些可能性极低但数量极多的标记偶然被抽中。

一个书里主动交代的简化

书里在生成函数的注释里写明:这是「效率较低的采样版本,没有使用键值缓存」21

也就是说,书里这段代码每生成一个字,都要把前面所有字重算一遍。 真实的推理框架不会这么做——这正是第 5 节讲的 KV 缓存要解决的问题。

判断(我们的,不是书里的): 这个简化是对的教学取舍,但它留下了一个理解上的坑: 读者会以为「分组查询注意力省显存」和「生成」是两件不相干的事。 实际上它们是同一件事的两面——分组查询省的正是这段代码里被省略掉的那个缓存。 书里两处相隔几十页,中间没有一句话把它们连起来。 如果错,会错在: 如果分组查询注意力在训练阶段也有显著的显存收益(它确实有,只是小得多), 那么把它完全归因于推理阶段就是片面的。

8. 四处替换的总账

这是这一章最该带走的一张表。

零件2017 年原版今天换来什么代价
归一化层归一化RMSNorm少算一个统计量、少一个参数,每层省 7%–64% 的归一化耗时几乎没有
位置编码正余弦,加到向量上旋转位置编码相对位置由数学结构保证;能外推到更长实现更复杂
注意力每个头一组键值分组共用键值推理时显存带宽(单位时间能从显存里搬多少数据)的压力大降质量略降
前馈层两层 + ReLU三层 + 门控质量更好没有机制解释,纯经验

判断(我们的,不是书里的):这四处替换的性质完全不同,值得分开看。 第一处是「删掉没用的」,第二处是「换一个更对的数学结构」, 第三处是「用质量换速度的明确交易」,第四处是「试出来更好」。 只有前两处有清楚的道理;第三处是明码标价的取舍;第四处是玄学。 一个模型架构里同时存在这四类零件,是这一行的常态。 如果错,会错在: 如果把第四处也当成有道理的设计去推广到别的结构上, 很可能得到一个更差的结果——经验结论不保证可迁移。

9. 作者的判断与证据

说法书里给了什么该怎么看
RMSNorm 有助于稳定学习只给了一句结论结论正确,但书没解释为什么可以省掉居中;我们在第 3 节补了出处
旋转位置编码「可以为注意力机制提供更强的上下文信息」只有这一句这句话说得太含糊——真正的机制是「相对位置由结构保证」,书没讲
分组查询「可以提高效率、节省显存」给了结论,没给数据方向正确,但书没说清省的是推理阶段的显存带宽
门控前馈层只给了代码和逐行注释书完全没讲为什么;原论文也没讲
每个模块都给了形状测试给了实际输出的张量形状这是好的工程习惯,能确认模块没写错

这一章的整体特点:书给的是「怎么写」,不是「为什么这么写」。 四个零件里,书对每一个都给了完整可读的代码和逐行注释, 但对「为什么换掉原来那个」几乎没有交代。 这正是「教程」和「教材」的差别——它保证你能跑起来,不保证你能判断该不该用。

10. 边界与局限

① 四篇原始论文,一篇都没引

RMSNorm、旋转位置编码、分组查询注意力、门控前馈层—— 这一章的四个核心零件各有明确的原始出处,书里一个都没提。 我们在脚注里逐一补上了核对过的一手来源。

② 关键概念缺失:KV 缓存

这是最影响理解的一处。 分组查询注意力存在的全部理由就是省 KV 缓存, 而「KV 缓存」这个词在书里只在生成函数的一句注释里出现过一次,正文从未解释。

结果是:读者知道「有 8 组键值头」这个事实,但不知道它省的是什么。

③ 没讲的取舍:书里的实现和 LLaMA2 官方并不完全一致

书里明说 LLaMA2 只有 700 亿参数版用了分组查询注意力,而它在小模型上也用了13。 这是个合理的教学选择,但它意味着这一章搭出来的不是「LLaMA2」,是「LLaMA2 风格」。

④ 完全没提的三件事

缺什么为什么重要
注意力的平方代价书里整章没有提序列长度带来的计算量增长,而这是所有长文本手段的根源(我们在第 5 节把它和键值缓存那笔线性账拆开补上了)
权重初始化的讲究代码里有初始化,但没解释为什么用那个分布、为什么按层数缩放
混合专家结构(MoE)今天大量开源模型把前馈层换成了多个专家 + 一个路由器,书里完全没有

⑤ 一处用词提醒

书里把旋转位置编码的两张表叫「实部」和「虚部」6。 这个叫法来自它的复数推导(第 03 章那段数学推导的延续), 但在代码里它们就是余弦表和正弦表。别被「虚部」这个词吓到。

11. 可带走的

  1. 今天的解码器和 2017 年的原版差四个零件,骨架完全一样,零件全换了;
  2. RMSNorm = 层归一化去掉「减平均值」这一步,只除以均方根;少一个参数、少一次统计;
  3. 它的论文贡献不是发明,是证明「少做一件事不会更差」——这类工作在工程上直接省钱;
  4. 旋转位置编码不给向量加东西,而是按位置转一个角度;
  5. 转角度的好处:两个位置的注意力分数天然只跟「隔多远」有关——相对位置由数学结构保证,不靠学;
  6. 旋转只作用在查询和键上,不作用在值上——位置是用来算相关性的,不该扭曲内容;
  7. 它能外推到训练时没见过的长度,因为它是一条公式而不是一张表;
  8. KV 缓存 = 把前面标记的键和值存起来重复用;生成时的瓶颈是搬运这个缓存,不是算力;
  9. 分组查询注意力 = 多个查询头共用一组键值头,省的正是这个缓存;质量略降是明码标价的代价;
  10. FlashAttention 改的是实现不是结构,结果完全一样,只是搬运更少;
  11. 门控前馈层 = 多一条通道当阀门,决定另一条通道每个数放多少过去;
  12. 门控是「白送」的:三条通道各占 2/3 宽,总参数量和原来差不多;
  13. 门控为什么有效,书没讲,原论文也没讲——用它,但别给它编原理;
  14. 温度是在「拉开」还是「抹平」分数之间的差距,不改变排名;
  15. top-k 是只留前 k 个候选,作用是堵住长尾;
  16. 四处替换的性质完全不同:删冗余、换结构、明码交易、纯经验——别一视同仁。

12. 原文地图

主题原书章原文位置
超参数配置5.1 动手实现一个 LLaMA2 大模型text/11-ch05-01-5-1-llama2.txt:15(搜「ModelConfig」) · text/11-ch05-01-5-1-llama2.txt:26(搜「n_layers: int = 12」) · text/11-ch05-01-5-1-llama2.txt:28(搜「n_kv_heads」) · text/11-ch05-01-5-1-llama2.txt:35(搜「flash_attn」)
RMSNorm5.1 动手实现一个 LLaMA2 大模型text/11-ch05-01-5-1-llama2.txt:58(搜「RMSNorm 可以⽤如下的数学公式」) · text/11-ch05-01-5-1-llama2.txt:64(搜「可学习的缩放参数」) · text/11-ch05-01-5-1-llama2.txt:70(搜「稳定学习过程」)
分组查询注意力与 repeat_kv5.1 动手实现一个 LLaMA2 大模型text/11-ch05-01-5-1-llama2.txt:108(搜「Grouped-Query Attention」) · text/11-ch05-01-5-1-llama2.txt:113(搜「repeat_kv」) · text/11-ch05-01-5-1-llama2.txt:143(搜「扩展和重塑张量」) · text/11-ch05-01-5-1-llama2.txt:291(搜「self.wk = nn.Linear」)
旋转位置编码5.1 动手实现一个 LLaMA2 大模型text/11-ch05-01-5-1-llama2.txt:154(搜「旋转嵌⼊」) · text/11-ch05-01-5-1-llama2.txt:160(搜「precompute_freqs_cis」) · text/11-ch05-01-5-1-llama2.txt:198(搜「实部和虚部」) · text/11-ch05-01-5-1-llama2.txt:221(搜「apply_rotary_emb」)
FlashAttention 开关5.1 动手实现一个 LLaMA2 大模型text/11-ch05-01-5-1-llama2.txt:303(搜「scaled_dot_product_attention」)
门控前馈层5.1 动手实现一个 LLaMA2 大模型text/11-ch05-01-5-1-llama2.txt:391(搜「输⼊维度的4倍」) · text/11-ch05-01-5-1-llama2.txt:392(搜「减少到2/3」) · text/11-ch05-01-5-1-llama2.txt:411(搜「F.silu」)
解码器层的组装5.1 动手实现一个 LLaMA2 大模型text/11-ch05-01-5-1-llama2.txt:469(搜「DecoderLayer 就是把我们上」)
生成函数、温度、top-k5.1 动手实现一个 LLaMA2 大模型text/11-ch05-01-5-1-llama2.txt:599(搜「def generate」) · text/11-ch05-01-5-1-llama2.txt:603(搜「没有使⽤键k/v cache」) · text/11-ch05-01-5-1-llama2.txt:615(搜「temperature == 0.0」) · text/11-ch05-01-5-1-llama2.txt:621(搜「top_k」)
参数量测试结果5.2 训练 Tokenizertext/12-ch05-02-5-2-tokenizer.txt:19(搜「82594560」)

Footnotes

  1. 出处:「5.1 动手实现一个 LLaMA2 大模型」第 3 段(text/11-ch05-01-5-1-llama2.txt:3,搜「2023年2⽉发布第⼀款」)。原文:Meta 于 2023 年 2 月发布第一款基于 Transformer 结构的大型语言模型 LLaMA,同年 7 月发布 LLaMA2。LLaMA2 的一手来源见第 04 章脚注。

  2. 出处:「5.1 动手实现一个 LLaMA2 大模型」第 15 段(text/11-ch05-01-5-1-llama2.txt:15,搜「ModelConfig」)、第 26 段(text/11-ch05-01-5-1-llama2.txt:26,搜「n_layers: int = 12」)、第 28 段(text/11-ch05-01-5-1-llama2.txt:28,搜「n_kv_heads」)、第 33 段(text/11-ch05-01-5-1-llama2.txt:33,搜「max_seq_len: int = 512」)。这份配置是照着 transformers 库的参数定义改出来的(编程里管这叫「继承」:沿用现成的一份定义,只写不一样的那部分),书里给的理由是方便后续导出为 Hugging Face 格式的模型。

  3. 出处:「5.2 训练 Tokenizer」第 19 段(text/12-ch05-02-5-2-tokenizer.txt:19,搜「82594560」)。这是把整个模型搭完之后跑测试打印出来的参数总数。

  4. 出处:「5.1 动手实现一个 LLaMA2 大模型」第 58 段(text/11-ch05-01-5-1-llama2.txt:58,搜「RMSNorm 可以⽤如下的数学公式」)、第 64 段(text/11-ch05-01-5-1-llama2.txt:64,搜「可学习的缩放参数」)、第 70 段(text/11-ch05-01-5-1-llama2.txt:70,搜「稳定学习过程」)。同一段说明也出现在讲 T5 的那一节(见第 04 章)。

  5. 补充(不在书里):《Root Mean Square Layer Normalization》,作者 Biao Zhang 与 Rico Sennrich,首次公开于 2019 年 10 月 16 日。摘要的原话是:他们「假设层归一化里的重新居中不变性是可有可无的」,去掉之后 RMSNorm「取得了与层归一化相当的表现,但在不同模型上把运行时间减少了 7% 到 64%」。来源:https://arxiv.org/abs/1910.07467(查阅于 2026-08-25)。

  6. 出处:「5.1 动手实现一个 LLaMA2 大模型」第 154 段(text/11-ch05-01-5-1-llama2.txt:154,搜「旋转嵌⼊」)、第 200 段(text/11-ch05-01-5-1-llama2.txt:200,搜「reshape_for_broadcast」)、第 221 段(text/11-ch05-01-5-1-llama2.txt:221,搜「apply_rotary_emb」)。「实部」「虚部」的叫法见第 198 段(text/11-ch05-01-5-1-llama2.txt:198,搜「实部和虚部」)。 2

  7. 出处:「5.1 动手实现一个 LLaMA2 大模型」第 160 段(text/11-ch05-01-5-1-llama2.txt:160,搜「precompute_freqs_cis」)与第 175 至 196 段的逐步说明(text/11-ch05-01-5-1-llama2.txt:189,搜「torch.outer」)。这段代码默认用的底数 theta 就是 10000,和第 03 章那个正余弦公式一致。

  8. 出处:「5.1 动手实现一个 LLaMA2 大模型」第 326 段(text/11-ch05-01-5-1-llama2.txt:326,搜「apply_rotary_emb(xq, xk」)。代码里传进旋转函数的只有查询和键两个张量。「为什么不转值」这层解释是我们补的(补充,不在书里,来自通用知识)。

  9. 出处:「第四章 大语言模型」第 152 段(text/10-ch04.txt:152,搜「RoPE」)与第 154 段(text/10-ch04.txt:154,搜「InternLM」)。

  10. 补充(不在书里):《RoFormer: Enhanced Transformer with Rotary Position Embedding》,第一作者 Jianlin Su,首次公开于 2021 年 4 月 20 日。摘要点明的三个性质是:用旋转矩阵编码绝对位置、在自注意力公式里自带相对位置依赖、以及「随相对距离增大而衰减的词间依赖」。来源:https://arxiv.org/abs/2104.09864(查阅于 2026-08-25)。

  11. 出处:「5.1 动手实现一个 LLaMA2 大模型」第 291 段(text/11-ch05-01-5-1-llama2.txt:291,搜「self.wk = nn.Linear」)。代码里查询的线性层输出宽度按 n_heads 算,键和值的按 n_kv_heads 算——这一处形状差异就是分组查询的全部实现。

  12. 出处:「5.1 动手实现一个 LLaMA2 大模型」第 113 段(text/11-ch05-01-5-1-llama2.txt:113,搜「repeat_kv」)与第 143 段(text/11-ch05-01-5-1-llama2.txt:143,搜「扩展和重塑张量」)。书里对这个函数给了四步逐行说明。

  13. 出处:「5.1 动手实现一个 LLaMA2 大模型」第 108 段(text/11-ch05-01-5-1-llama2.txt:108,搜「Grouped-Query Attention」)与第 109 段(text/11-ch05-01-5-1-llama2.txt:109,搜「提⾼模型的效率」)。 2

  14. 补充(不在书里):《GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints》,第一作者 Joshua Ainslie,首次公开于 2023 年 5 月 22 日。摘要指出多查询注意力「大幅加快了解码器推理,但可能损害质量」;分组查询用「多于一个、少于查询头数」的键值头做折中,结果是「质量接近多头注意力,速度接近多查询注意力」;并给出了用原始预训练算力 5% 把已有多头检查点转换过去的做法。来源:https://arxiv.org/abs/2305.13245(查阅于 2026-08-25)。多查询注意力的原始出处是《Fast Transformer Decoding: One Write-Head is All You Need》,作者 Noam Shazeer,首次公开于 2019 年 11 月 6 日,摘要说明它让「键和值在所有注意力头之间共享」。来源:https://arxiv.org/abs/1911.02150(查阅于 2026-08-25)。 2

  15. 出处:「5.1 动手实现一个 LLaMA2 大模型」第 303 段(text/11-ch05-01-5-1-llama2.txt:303,搜「scaled_dot_product_attention」)与第 35 段(text/11-ch05-01-5-1-llama2.txt:35,搜「flash_attn」)。代码通过检测 PyTorch 有没有这个函数来决定走哪条路,走不通就退回手写的注意力实现。

  16. 补充(不在书里):《FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness》,第一作者 Tri Dao,首次公开于 2022 年 5 月 27 日。它的论点是既有的近似注意力方法用模型质量换复杂度、却往往拿不到真实的墙钟加速,缺的原则是让注意力算法「对读写敏感」;论文报告的加速包括 BERT-large 上 15%、GPT-2 上 3 倍。来源:https://arxiv.org/abs/2205.14135(查阅于 2026-08-25)。

  17. 出处:「5.1 动手实现一个 LLaMA2 大模型」第 411 段(text/11-ch05-01-5-1-llama2.txt:411,搜「F.silu」)与第 413 段(text/11-ch05-01-5-1-llama2.txt:413,搜「SILU 激活函数」)。原文对前向计算的描述:输入先过第一层线性变换和 SiLU 激活,结果乘以输入过第三层线性变换的结果,最后过第二层线性变换和 dropout。

  18. 出处:「5.1 动手实现一个 LLaMA2 大模型」第 391 段(text/11-ch05-01-5-1-llama2.txt:391,搜「输⼊维度的4倍」)与第 392 段(text/11-ch05-01-5-1-llama2.txt:392,搜「减少到2/3」)。书里的注释只说「设置为输入维度的 4 倍,然后减少到 2/3,最后确保是 multiple_of 的倍数」,没有解释为什么;「因为多了一条通道要把参数量拉回来」这层解释是我们补的(补充,不在书里,来自通用知识)。

  19. 补充(不在书里):《GLU Variants Improve Transformer》,作者 Noam Shazeer,首次公开于 2020 年 2 月 12 日。摘要说门控线性单元是「两路线性投影的逐位乘积,其中一路先过一个非线性函数」;论文把 Transformer 前馈层里的激活换成各种变体做测试,发现某些变体优于标准的 ReLU 或 GELU。摘要里没有任何关于「为什么有效」的解释,只报告了实验质量提升。来源:https://arxiv.org/abs/2002.05202(查阅于 2026-08-25)。

  20. 出处:「5.1 动手实现一个 LLaMA2 大模型」第 469 段(text/11-ch05-01-5-1-llama2.txt:469,搜「DecoderLayer 就是把我们上」)。

  21. 出处:「5.1 动手实现一个 LLaMA2 大模型」第 599 段(text/11-ch05-01-5-1-llama2.txt:599,搜「def generate」)与第 603 段(text/11-ch05-01-5-1-llama2.txt:603,搜「没有使⽤键k/v cache」)。原文注释:「效率较低的采样版本,没有使用键 k/v cache」。 2

  22. 出处:「5.1 动手实现一个 LLaMA2 大模型」第 615 段(text/11-ch05-01-5-1-llama2.txt:615,搜「temperature == 0.0」)、第 620 段(text/11-ch05-01-5-1-llama2.txt:620,搜「logits = logits / temperature」)、第 621 段(text/11-ch05-01-5-1-llama2.txt:621,搜「top_k」)。书里只给了代码和一句注释,「除以温度为什么能控制随机程度」和「top-k 堵长尾」这两层解释是我们补的(补充,不在书里,来自通用知识)。 2