跳到主要内容

预训练实战 — 一笔能自己验算的训练账

这一章讲三件事: 预训练阶段模型到底在解什么题(不止「预测下一个词」); 训练前人工定好的那些设置——行话叫超参数——与并行方案各有什么成型配方 (为什么几千张卡能朝一个目标使劲),以及全书最硬的一节——参数、算力、时间、 显存四笔账,全部可复算。

超参数里最要紧的两个,先交代掉:学习率——每一步参数挪多大;

另一个是每步喂多少 token——这个量行话单叫一个字:批——后文表里 GPT-3 那行 「32K→3.2M」说的就是它。 主走查: 拿 LLaMA-7B/65B 的真实数字,把四笔账逐笔算一遍,最后一笔和论文实测对上。

1. 预训练任务:三道题,各练一种本事

第 01 章说预训练是「自学习」,正式名字叫自监督:不给人工标注, 让模型从文本自身的结构里出题自己答。原书列了三道题1

第一题:语言建模(LM),就是「预测下一个词」。 为什么这么简单的题能练出理解? 原书摘录了 OpenAI 前首席科学家 Sutskever 的一段访谈:读侦探小说读到最后一页, 侦探说「我现在揭晓凶手,这个人就是__」——要把这个词预测准,前面几百页的伏笔你都得真读懂; 预测得越来越准,对文本的理解就越来越深2。这道题还有两个变体: 前缀语言建模(随机切一刀,只拿后半截的损失来训练,适合「输入-输出」式的任务, 但同样数据下略差于标准版)3;FIM 填中(把一段话的前缀和后缀拼在一起, 让模型补出中间那段——代码补全就是「给上下文填中间」,所以代码模型普遍加练这道题)4

第二题:去噪自编码(DAE)。 把原文随机删改一番,让模型还原——BERT、T5 靠它训练。 它更擅长「理解」而不是「续写」,所以在大模型时代用得少,代表只剩 T5 和 GLM-130B5

第三题:混合去噪器(UL2)。 把前两题统一成一族去噪任务: S 档就是续写,R 档小删小补(遮约 15%、每段 3~5 个 token), X 档大删大补(段长超 12 或遮约 50%);训练时在输入开头放一个 [S]/[R]/[X] 标记(特殊符号), 告诉模型这回合按哪档答题。UL2 和 PaLM-2 用的这种6

2. 优化设置:一张有出处的配方表

2.1 批量与学习率

(这两个词章首已交代:批,就是每步喂多少 token——它的正式名字叫批量;学习率,就是每一步参数挪多大。)

原书表 6.1 汇总了二十来个模型的真实超参,挑几行就能看出配方的形状7:

模型批量(每步喂多少 token)学习率(峰值)优化器
GPT-332K → 3.2M(逐步加大)6×10⁻⁵Adam
PaLM1M → 4M1×10⁻²(逆平方根衰减)Adafactor
LLaMA-24M1.5×10⁻⁴AdamW
DeepSeek18M3.2×10⁻⁴AdamW

批量(batch size)普遍是百万级,而且流行先小后大: 开局小批量让损失快速下降,后期大批量稳步收敛——GPT-3 从 32K 一路加到 3.2M8学习率的用法是先预热(warmup:从接近 0 线性升到峰值, 占全部步数的 0.10.5%——开局参数是随机的,大步走容易翻车), 升到峰值(常见 5×10⁻⁵1×10⁻⁴)后再缓缓衰减,收尾时大约剩峰值的一成9

2.2 优化器与稳定四件套

优化器(决定「每一步参数具体怎么挪」的算法)几乎清一色是 Adam 家族—— Adam 自己就是一种优化器的名字。 Adam 的两个招数:给梯度带上惯性(加权平均历史梯度,不被单批数据的噪声(数据里的随机杂质)带偏); 按各参数近期的波动自动调节步长。它的三个超参在大模型圈有统一口径: β1=0.9、β2=0.95、ε=10⁻⁸——注意 β2 不是教科书默认值 0.999, 这个改动是大模型训练的实战经验10。PaLM 和 T5 用更省显存的变体 Adafactor10

训练大模型最常见的意外是「损失突然飙升」,原书给了四件稳定器11:

  • 梯度裁剪:梯度(每个参数该往哪挪、挪多少的那串数)偶尔爆炸, 超过阈值就砍回去,大模型训练的阈值通行 1.0;
  • 定时存档:每隔一段存一次模型,损失异常飙升时回滚到上一个存档、跳过可疑数据;
  • 权重衰减:每步更新时顺手把参数往零拉一点(系数 0.1),防参数无节制膨胀;
  • dropout(随机关掉一些神经元防过拟合的老办法)基本不用: 在海量数据+归一化结构面前没必要——过拟合不是大模型预训练的主要矛盾。

3. 可扩展训练:三种切法,各治一种瓶颈

一张卡的显存装不下模型,一张卡的算力等不起训练——所以要多卡。 怎么切,原书讲了三种,各治一种瓶颈12(张量:多维数组的总称—— 向量是一维、矩阵是二维,模型的参数就是一堆矩阵,「张量并行」就是把这些矩阵切块):

切法怎么切治什么留下的问题
数据并行每卡复制整个模型,各喂不同的数据,最后把梯度求平均吞吐量显存浪费:每卡都存一份全量
流水线并行按层切:第 1-2 层在卡 1,第 3-4 层在卡 2单层装不下有「气泡」:卡 1 要等卡 2 算完
张量并行把每层的参数矩阵横竖切块,分到多卡同算单层都算不动卡间通信量大

流水线并行的「气泡」有解法:梯度累积——算完一批不急着更新参数, 多攒几批再一起更新;这样卡 1 算完第一批的前向,不用等反向回来,接着算第二批, 气泡就被填上了13。(「前向」指从输入算出预测,「反向」指从损失回推每个参数的梯度—— 一正一反合起来才是一步训练,反向的算力开销约为前向的两倍14。)

在这三种切法之上,DeepSpeed 的 ZeRO 专治数据并行的显存浪费: 不让你每卡存全量,而是把优化器状态、梯度、参数分三档逐层切开摊到各卡,用时再取15。 还有两件省显存的标配:激活重算——前向时只存每层的输入, 反向要用中间结果时现场重算一遍,拿算力换显存16;混合精度—— 前向反向用 16 位浮点(省一半显存、快一倍),但留一份 32 位主参数做最终更新。 16 位格式有两种:FP16 范围小(±65,504);bfloat16——尾数位数更多、范围大(约 10³⁸)不易溢出——新一代模型全用它17

4. 主走查:LLaMA 的四笔账,逐笔算给你看

这是原书 6.4 节,全书工程含量最高的几页。四笔账全部给出公式,我们逐笔代真数复算。

第 ① 笔:参数账——6,738,415,616 是怎么数出来的

第 05 章的零件清单一换算就是参数公式。一个词表 V、L 层、宽度 H、FFN 中间宽 H′ 的模型:

参数总数 = 2VH + H + L·(4H² + 3HH′ + 2H)
│ │ │ └ 每层 2 个归一化(2H)
│ │ └ FFN 三矩阵(2HH′+HH′)
│ └ 注意力四矩阵(Q/K/V/O 各 H²)
└ 输入嵌入 VH + 输出映射 VH

代入 LLaMA-7B:V=32,000,L=32,H=4096,H′=11,008
= 2×32000×4096 + 4096 + 32×(4×4096² + 3×4096×11008 + 2×4096)
= 6,738,415,616
图说:算出来 67.38 亿,与 LLaMA-7B 的真实参数量分毫不差。

第 ② 笔:算力账——为什么总成本 ≈ 6CP

浮点运算(一次小数加减乘除;「浮点」是计算机表示小数的方式)次数(FLOPs)是训练成本的度量。两块大头: 矩阵乘法有精确公式(n×m 矩阵乘 m×p 矩阵 = 2nmp 次)18; 注意力那部分约为线性变换的 T/(6H)——当序列长 T 小于隐层宽 H 时, 占比不足六分之一,可以忽略;而线性变换的参数占全模型 95% 以上。 于是总成本收敛成一个漂亮近似:训练成本 ≈ 6 × C × P(C=训练 token 数,P=参数量; 6 = 前向 2 + 反向 4);开激活重算多一遍前向,变成 8CP19。 代进真数:LLaMA-7B 训 10 亿 token ≈ 6 × 6.74×10⁹ × 10⁹ ≈ 4.04×10¹⁹ 次浮点运算20

第 ③ 笔:时间账——20.6 天,论文说 21 天

已知(全是真数):LLaMA-65B,P = 6.5×10¹⁰,C = 1.4×10¹² token,开了激活重算
总算力 = 8CP = 8 × 1.4×10¹² × 6.5×10¹⁰ ≈ 7.28×10²³ 次
硬件:2048 张 A100,单张峰值 3.12×10¹⁴ 次/秒;
实际只能跑到峰值的 30~70%,按 2×10¹⁴ 次/秒估
时间 = 7.28×10²³ ÷ (2048 × 2×10¹⁴) ≈ 1.78×10⁶ 秒 ≈ 20.6 天
论文实际报告:21 天。——估算闭环。
图说:这就是第 03 章表 3.1 里「2048 卡 21 天」那行的来历;
任何一个模型的训练天数,你都能这样先算出来再对照。

第 ④ 笔:显存账——16P 字节这条铁律从哪来

混合精度训练下,每张卡要存:16 位参数 2P 字节 + 16 位梯度 2P 字节 + 优化器的三份 32 位状态(参数、动量、方差)各 4P 字节——合计 16P 字节。 这就是「全参数训练至少要 16 倍参数量的显存」的出处。ZeRO 三档逐级往下砍: ZeRO-1 切优化器状态 → 4P + 12P/N;ZeRO-2 再切梯度 → 2P + 14P/N; ZeRO-3 连参数也切 → 16P/N(N = 数据并行卡数)21

剩下的零头原书也算了:激活值(反向时要用的中间结果,开激活重算后约占几 GB)、 softmax 中间量(约为输入的 2 倍)、PyTorch 内核(0.81GB)、 ZeRO 自身(14GB)、显存碎片(0.5~1GB)。全部代进 LLaMA-7B 双卡 A800-80G 的场景 (批量 8、开 FlashAttention+激活重算+ZeRO-3)22:

每卡:参数+优化器 50.20 GB(ZeRO-3 切半后的 16P/2)
+ 激活 6.20 GB + softmax 中间量 3.91 GB + 其他 6 GB
≈ 66 GB —— 两张 80GB 的卡,刚好装下,还剩约 14GB 余量。

5. 经验法则:动手前先算这三笔

把 §4 的公式倒过来用,就是原书给的实战口诀23:

  • 全参数训练:显存 ≥ 16 × 参数量。 13B 模型 → 至少 208GB → 3 张 A800-80G 才够下限; 但 3 张卡时每张只剩约 10GB 给数据,批量只能开到 2,效率太低——建议 4 张,批量开 12; 同理 30B → 8 张,65B → 16 张(均按 80GB 卡)。
  • ≤30B 且卡少于 16 张:数据并行 + ZeRO 就够(DeepSpeed 一个 JSON(一种常见的配置文件格式)配好三档); 70B 级、上百上千张卡:必须上完整 3D 并行。
  • 多卡之间拼的是通信:单机内用 NVLink/NVSwitch 高速互联, 跨机器要 InfiniBand(一种专为机间高速数据传输设计的网络)—— 否则算力会被等待拖垮24

(对照:复旦版教材讲分布式训练(把一次训练拆到成百上千张卡上一起跑)时铺开「三堵墙+三种切法+集群拓扑」的体系; 本书的打法是给你一支笔——公式、代数、与论文实测对账。两本对照读, 一个知道系统全貌,一个会算自己的账。)

判断(我们的,不是书里的): 6.4 节是这本书区别于一般综述的地方—— 它把「训练一个模型要多少资源」从玄学变成了每个读者都能复算的算术题, 并且主动拿论文实测(21 天)来验证自己的估算(20.6 天)。 这种「公式 → 代数 → 实测闭环」的写法,是工程教材而非研究综述的笔法。 如果错,会错在: 公式里的近似(忽略注意力、按 6CP/8CP 估)在序列特别长 (T 接近或超过 H)或 MoE 架构下会失真——那时注意力占比和激活参数都变了, 账要重算。LLaMA 这类稠密(所有参数每步都参与计算,与 MoE 相对)模型、T<H 的常规场景下,近似是稳的。

6. 作者的判断与证据

  • 可复算的硬证据: 参数公式与 LLaMA-7B 真实值分毫不差25; 时间估算 20.6 天对论文 21 天26;显存公式逐项有出处(式 6.11~6.15)22
  • 来自厂商与库的口径: GPU 实际利用率为峰值的 30~70% 是原书给的工程经验区间26; A100 峰值算力来自英伟达官网(原书脚注)。
  • 配方类数字都有出处: warmup 0.1~0.5%、β2=0.95、梯度裁剪 1.0、权重衰减 0.1, 出自原书对主流模型训练报告的汇总(表 6.1 同步可查)791011

7. 边界与局限

  • 6CP/8CP 的近似假设 T<H(序列比隐层短);长上下文训练(32K+)下注意力项占比上升,账要按式 6.6 重算。
  • 显存估算以 LLaMA 的稠密结构为准;MoE 模型(每步只激活一部分参数)套用时, 「16P」要换成「16 × 激活参数 + 路由与通信开销」,原书没有展开。
  • 时间估算只算了浮点运算,数据读写、多机同步的耗时被假设摊进「30~70% 利用率」里; 跨机网络差的时候,这个区间会跌破下限。
  • 表 6.1 的配方止于 2024 年初的模型;β2=0.95 这类「统一口径」未来仍可能变。

8. 可带走的

  1. 预训练三道题:下一词预测(主线)、去噪还原(理解向)、混合去噪器(统一族);FIM 变体专供代码补全;
  2. 侦探小说测验:能把「凶手是__」预测准,等于前面几百页都读懂了——下一词预测为什么有效的最直白解释;
  3. 配方:批量百万级、先小后大;学习率先预热(0.1~0.5% 步数)再衰减到一成;Adam 的 β2 要改成 0.95;
  4. 稳定四件套:梯度裁剪 1.0、定时存档、权重衰减 0.1、不用 dropout;
  5. 三种切法:数据并行切数据(废显存)、流水线切层(有气泡)、张量切矩阵(费通信);ZeRO 三档把 16P 砍到 16P/N;
  6. 参数公式 2VH+H+L(4H²+3HH′+2H):LLaMA-7B = 6,738,415,616,与真实值分毫不差;
  7. 训练成本 ≈ 6CP(开激活重算 8CP);LLaMA-65B:7.28×10²³ 次 ÷ 2048 卡 ≈ 20.6 天(论文 21 天);
  8. 显存铁律 16P 字节:13B→208GB→建议 4 卡;30B→8 卡;65B→16 卡;
  9. ≤30B、少于 16 卡:数据并行+ZeRO 就够;更大规模上 3D 并行+NVLink/InfiniBand;
  10. 估算先行的习惯:任何训练计划,先按 §4 的四个公式算一遍再开机。

9. 原文地图

主题原书章原文位置
三个预训练任务总览6.1text/30-ch06-6-model-pre-training.txt:51(搜「three common pre-training tasks」)
Sutskever 侦探小说例例 6.1text/30-ch06-6-model-pre-training.txt:65(搜「detective novel」)
前缀 LM 与 FIM6.1.1text/30-ch06-6-model-pre-training.txt:75(搜「slightly inferior」) · text/30-ch06-6-model-pre-training.txt:77(搜「fill-in-the-middle」)
去噪自编码6.1.2text/30-ch06-6-model-pre-training.txt:91(搜「T5 and GLM-130B」)
UL2 三档6.1.3text/30-ch06-6-model-pre-training.txt:97(搜「15% of tokens」)
表 6.1 各模型配方表 6.1text/30-ch06-6-model-pre-training.txt:121(搜「3.2M」) · text/30-ch06-6-model-pre-training.txt:359(搜「18M」)
批量先小后大6.2.1text/30-ch06-6-model-pre-training.txt:429(搜「32K tokens to 3.2M」)
学习率预热与衰减6.2.2text/30-ch06-6-model-pre-training.txt:433(搜「0.1–0.5%」)
Adam 与 β2=0.956.2.3text/30-ch06-6-model-pre-training.txt:439(搜「β2 = 0.95」)
稳定四件套6.2.4text/30-ch06-6-model-pre-training.txt:445(搜「clipping threshold」) · text/30-ch06-6-model-pre-training.txt:451(搜「dropout is rarely used」)
3D 并行6.3.1text/30-ch06-6-model-pre-training.txt:463(搜「Data parallelism」) · text/30-ch06-6-model-pre-training.txt:465(搜「gradient accumulation」) · text/30-ch06-6-model-pre-training.txt:467(搜「split into two submatrices」)
ZeRO 与 FSDP6.3.2text/30-ch06-6-model-pre-training.txt:471(搜「zero redundancy optimizer」)
激活重算与混合精度6.3.3-6.3.4text/30-ch06-6-model-pre-training.txt:475(搜「activation recomputation」) · text/30-ch06-6-model-pre-training.txt:479(搜「BF16」)
参数公式与 7B 复算6.4.1text/30-ch06-6-model-pre-training.txt:499(搜「2VH+H+L」) · text/30-ch06-6-model-pre-training.txt:503(搜「6,738,415,616」)
矩阵乘 2nmp6.4.2 Tipstext/30-ch06-6-model-pre-training.txt:513(搜「2nmp」)
6CP/8CP 推导6.4.2text/30-ch06-6-model-pre-training.txt:527(搜「6CP」)
7B 训 1B token 的 FLOPs6.4.2text/30-ch06-6-model-pre-training.txt:533(搜「4.04」)
65B 时间估算 20.6 天6.4.3text/30-ch06-6-model-pre-training.txt:541(搜「20.6 days」)
显存公式与 16P6.4.4text/30-ch06-6-model-pre-training.txt:551(搜「16P」)
ZeRO 三档显存6.4.4text/30-ch06-6-model-pre-training.txt:553(搜「ZeRO-1」) · text/30-ch06-6-model-pre-training.txt:557(搜「16P∕ND」)
7B 双卡 66GB6.4.4text/30-ch06-6-model-pre-training.txt:611(搜「66 GB」)
卡数经验法则6.4.4text/30-ch06-6-model-pre-training.txt:613(搜「at least 4 GPUs」)
3D 并行与硬件互联6.5text/30-ch06-6-model-pre-training.txt:625(搜「InfiniBand」)

Footnotes

  1. 出处:「6.1 Pre-training Tasks」第 51 段(text/30-ch06-6-model-pre-training.txt:51,搜「three common pre-training tasks」)。

  2. 出处:例 6.1 第 65-67 段(text/30-ch06-6-model-pre-training.txt:65,搜「detective novel」)。这是 Sutskever 访谈摘录(原书脚注给出 transcript 链接),是观点阐述而非实验结论。

  3. 出处:「6.1.1 Language Modeling」第 75 段(text/30-ch06-6-model-pre-training.txt:75,搜「slightly inferior」)。

  4. 出处:「6.1.1 Language Modeling」第 77-81 段(text/30-ch06-6-model-pre-training.txt:77,搜「fill-in-the-middle」;text/30-ch06-6-model-pre-training.txt:81,搜「code completion」)。

  5. 出处:「6.1.2 Denoising Autoencoding」第 91 段(text/30-ch06-6-model-pre-training.txt:91,搜「T5 and GLM-130B」)。

  6. 出处:「6.1.3 Mixture-of-Denoisers」第 97-99 段(text/30-ch06-6-model-pre-training.txt:97,搜「15% of tokens」;text/30-ch06-6-model-pre-training.txt:99,搜「[R], [S], [X]」)。

  7. 出处:表 6.1(text/30-ch06-6-model-pre-training.txt:103,搜「Table 6.1」)。GPT-3 行(text/30-ch06-6-model-pre-training.txt:121,搜「3.2M」)、DeepSeek 行(text/30-ch06-6-model-pre-training.txt:359,搜「18M」)。 2

  8. 出处:「6.2.1 Batch-Based Training」第 429 段(text/30-ch06-6-model-pre-training.txt:429,搜「32K tokens to 3.2M」)。

  9. 出处:「6.2.2 Learning Rate」第 433 段(text/30-ch06-6-model-pre-training.txt:433,搜「0.1–0.5%」)。 2

  10. 出处:「6.2.3 Optimizer」第 439 段(text/30-ch06-6-model-pre-training.txt:439,搜「β2 = 0.95」)。Adam 的 β2 默认 0.999 是通用知识,大模型训练改用 0.95 是原书明确给的口径。 2 3

  11. 出处:「6.2.4 Stable Optimization Techniques」第 445-451 段(text/30-ch06-6-model-pre-training.txt:445,搜「clipping threshold」;text/30-ch06-6-model-pre-training.txt:451,搜「dropout is rarely used」)。 2

  12. 出处:「6.3.1 3D Parallel Training」第 463-467 段(text/30-ch06-6-model-pre-training.txt:463,搜「Data parallelism」;text/30-ch06-6-model-pre-training.txt:467,搜「split into two submatrices」)。

  13. 出处:「6.3.1 3D Parallel Training」第 465 段(text/30-ch06-6-model-pre-training.txt:465,搜「gradient accumulation」)。

  14. 出处:「6.4.2 Estimation of Training Cost」脚注 2(text/30-ch06-6-model-pre-training.txt:689,搜「roughly double the computational overhead」)。反向传播算力约为前向两倍,这是 6CP 公式里 2+4=6 的来源。

  15. 出处:「6.3.2 Zero Redundancy Optimizer」第 471 段(text/30-ch06-6-model-pre-training.txt:471,搜「zero redundancy optimizer」)。

  16. 出处:「6.3.3 Activation Recomputation」第 475 段(text/30-ch06-6-model-pre-training.txt:475,搜「activation recomputation」)。

  17. 出处:「6.3.4 Mixed-Precision Training」第 479 段(text/30-ch06-6-model-pre-training.txt:479,搜「BF16」)。

  18. 出处:「6.4.2 Tips(Computational Cost of Matrix Multiplication)」第 513 段(text/30-ch06-6-model-pre-training.txt:513,搜「2nmp」)。

  19. 出处:「6.4.2 Estimation of Training Cost」第 527-531 段(text/30-ch06-6-model-pre-training.txt:527,搜「6CP」)。注意力成本约为线性变换的 T/(6H);线性参数占 95% 以上,故 P 可直接代入。

  20. 出处:「6.4.2 Estimation of Training Cost」第 533 段(text/30-ch06-6-model-pre-training.txt:533,搜「4.04」)。

  21. 出处:「6.4.4 Training Memory Estimation」第 551-557 段(text/30-ch06-6-model-pre-training.txt:551,搜「16P」;text/30-ch06-6-model-pre-training.txt:557,搜「16P∕ND」)。

  22. 出处:「6.4.4 Example: Estimation of Total Memory Consumption」第 611 段(text/30-ch06-6-model-pre-training.txt:611,搜「66 GB」)。各分项:参数与优化器 50.20GB、激活 6.20GB、中间结果 3.91GB、其他约 6GB。 2

  23. 出处:「6.4.4 Training Memory Estimation」第 613 段(text/30-ch06-6-model-pre-training.txt:613,搜「at least 4 GPUs」)。

  24. 出处:「6.5 Pre-training Code Practice」第 625 段(text/30-ch06-6-model-pre-training.txt:625,搜「InfiniBand」)。

  25. 出处:「6.4.1 Parameter Calculation」第 499-505 段(text/30-ch06-6-model-pre-training.txt:499,搜「2VH+H+L」;text/30-ch06-6-model-pre-training.txt:503,搜「6,738,415,616」;text/30-ch06-6-model-pre-training.txt:505,搜「exactly the same」)。

  26. 出处:「6.4.3 Training Time Estimation」第 541 段(text/30-ch06-6-model-pre-training.txt:541,搜「20.6 days」)。30~70% 的利用率区间与「论文报 21 天」均出此段。 2