预训练实战 — 一笔能自己验算的训练账
这一章讲三件事: 预训练阶段模型到底在解什么题(不止「预测下一个词」); 训 练前人工定好的那些设置——行话叫超参数——与并行方案各有什么成型配方 (为什么几千张卡能朝一个目标使劲),以及全书最硬的一节——参数、算力、时间、 显存四笔账,全部可复算。
超参数里最要紧的两个,先交代掉:学习率——每一步参数挪多大;
另一个是每步喂多少 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-3 | 32K → 3.2M(逐步加大) | 6×10⁻⁵ | Adam |
| PaLM | 1M → 4M | 1×10⁻²(逆平方根衰减) | Adafactor |
| LLaMA-2 | 4M | 1.5×10⁻⁴ | AdamW |
| DeepSeek | 18M | 3.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. 可带走的
- 预训练三道题:下一词预测(主线)、去噪还原(理解向)、混合去噪器(统一族);FIM 变体专供代码补全;
- 侦探小说测验:能把「凶手是__」预测准,等于前面几百页都读懂了——下一词预测为什么有效的最直白解释;
- 配方:批量百万级、先小后大;学习率先预热(0.1~0.5% 步数)再衰减到一成;Adam 的 β2 要改成 0.95;
- 稳定四件套:梯度裁剪 1.0、定时存档、权重衰减 0.1、不用 dropout;
- 三种切法:数据并行切数据(废显存)、流水线切层(有气泡)、张量切矩阵(费通信);ZeRO 三档把 16P 砍到 16P/N;
- 参数公式 2VH+H+L(4H²+3HH′+2H):LLaMA-7B = 6,738,415,616,与真实值分毫不差;
- 训练成本 ≈ 6CP(开激活重算 8CP);LLaMA-65B:7.28×10²³ 次 ÷ 2048 卡 ≈ 20.6 天(论文 21 天);
- 显存铁律 16P 字节:13B→208GB→建议 4 卡;30B→8 卡;65B→16 卡;
- ≤30B、少于 16 卡:数据并行+ZeRO 就够;更大规模上 3D 并行+NVLink/InfiniBand;
- 估算先行的习惯:任何训练计划,先按 §4 的四个公式算一遍再开机。
9. 原文地图
| 主题 | 原书章 | 原文位置 |
|---|---|---|
| 三个预训练任务总览 | 6.1 | text/30-ch06-6-model-pre-training.txt:51(搜「three common pre-training tasks」) |
| Sutskever 侦探小说例 | 例 6.1 | text/30-ch06-6-model-pre-training.txt:65(搜「detective novel」) |
| 前缀 LM 与 FIM | 6.1.1 | text/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.2 | text/30-ch06-6-model-pre-training.txt:91(搜「T5 and GLM-130B」) |
| UL2 三档 | 6.1.3 | text/30-ch06-6-model-pre-training.txt:97(搜「15% of tokens」) |
| 表 6.1 各模型配方 | 表 6.1 | text/30-ch06-6-model-pre-training.txt:121(搜「3.2M」) · text/30-ch06-6-model-pre-training.txt:359(搜「18M」) |
| 批量先小后大 | 6.2.1 | text/30-ch06-6-model-pre-training.txt:429(搜「32K tokens to 3.2M」) |
| 学习率预热与衰减 | 6.2.2 | text/30-ch06-6-model-pre-training.txt:433(搜「0.1–0.5%」) |
| Adam 与 β2=0.95 | 6.2.3 | text/30-ch06-6-model-pre-training.txt:439(搜「β2 = 0.95」) |
| 稳定四件套 | 6.2.4 | text/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.1 | text/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 与 FSDP | 6.3.2 | text/30-ch06-6-model-pre-training.txt:471(搜「zero redundancy optimizer」) |
| 激活重算与混合精度 | 6.3.3-6.3.4 | text/30-ch06-6-model-pre-training.txt:475(搜「activation recomputation」) · text/30-ch06-6-model-pre-training.txt:479(搜「BF16」) |
| 参数公式与 7B 复算 | 6.4.1 | 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」) |
| 矩阵乘 2nmp | 6.4.2 Tips | text/30-ch06-6-model-pre-training.txt:513(搜「2nmp」) |
| 6CP/8CP 推导 | 6.4.2 | text/30-ch06-6-model-pre-training.txt:527(搜「6CP」) |
| 7B 训 1B token 的 FLOPs | 6.4.2 | text/30-ch06-6-model-pre-training.txt:533(搜「4.04」) |
| 65B 时间估算 20.6 天 | 6.4.3 | text/30-ch06-6-model-pre-training.txt:541(搜「20.6 days」) |
| 显存公式与 16P | 6.4.4 | text/30-ch06-6-model-pre-training.txt:551(搜「16P」) |
| ZeRO 三档显存 | 6.4.4 | text/30-ch06-6-model-pre-training.txt:553(搜「ZeRO-1」) · text/30-ch06-6-model-pre-training.txt:557(搜「16P∕ND」) |
| 7B 双卡 66GB | 6.4.4 | text/30-ch06-6-model-pre-training.txt:611(搜「66 GB」) |
| 卡数经验法则 | 6.4.4 | text/30-ch06-6-model-pre-training.txt:613(搜「at least 4 GPUs」) |
| 3D 并行与硬件互联 | 6.5 | text/30-ch06-6-model-pre-training.txt:625(搜「InfiniBand」) |