8. 预训练工程
- 拆解一个训练 step 的计算与内存
- 区分 batch、累积和并行策略
- 诊断 NaN、OOM、吞吐下降和恢复失败
- 完成 TinyGPT 训练
- 理解反向传播和优化器
本章产物一份训练容量估算与故障处置 Runbook。
8.1 预训练目标
对 token 序列 x_1 ... x_T,因果语言模型最大化:
P(x_1, ..., x_T)
= Π_t P(x_t | x_1, ..., x_(t-1))取负对数后得到可相加的 loss:
Loss = -Σ_t log P(x_t | x_<t)这个目标不直接告诉模型“什么是事实”或“如何推理”,而是要求它压缩和预测数据中的规律。为了降低大量不同文本上的预测误差,模型会形成语法、语义、事实共现和任务模式等内部表示。
8.2 一个训练 step 发生了什么
1. DataLoader 取 batch
2. tokenizer/数据文件提供 input_ids
3. forward 产生 logits
4. 交叉熵产生 loss
5. backward 产生所有可训练参数的梯度
6. gradient clipping(可选)
7. optimizer 更新权重与状态
8. scheduler 更新学习率
9. 记录指标、定期验证和保存 checkpoint训练通常比推理更吃内存,因为不仅需要权重,还需要:
- 激活值,用于 backward。
- 梯度。
- 优化器状态,例如 Adam 的一阶和二阶矩。
- 高精度主权重或 loss scaling 相关状态。
- 通信和临时缓冲区。
FP16 权重约 2 字节/参数,但使用 AdamW 的完整训练总占用可能远高于 2 字节/参数。
8.3 Batch、Micro-batch 和梯度累积
显存放不下目标 batch 时:
有效 batch
= 每设备 micro-batch
× gradient_accumulation_steps
× 数据并行设备数梯度累积在多次 micro-batch 后才执行 optimizer step。若 loss 默认取均值,框架通常会正确缩放;手写训练时要注意除以累积步数。
梯度累积能减少峰值内存,但不能完全等价于大 batch:
- BatchNorm 会不同,不过 LLM 常不用 BatchNorm。
- dropout 随机性不同。
- 每个 optimizer step 前的数据组织不同。
- 通信和调度效率不同。
8.4 混合精度
- FP32:范围和精度较高,内存与带宽开销大。
- FP16:内存小、硬件快,但指数范围较窄,容易 overflow/underflow。
- BF16:指数范围接近 FP32,尾数精度较低,常用于训练。
混合精度并不是把所有东西都盲目转成低精度。常见做法是矩阵乘法使用低精度,某些累加、归一化或优化器状态保留更高精度。
Apple MPS 与 NVIDIA CUDA 支持细节不同。不要把 CUDA 教程中的 bitsandbytes、torch.cuda.amp 参数原样套到 Mac。
8.5 Gradient Checkpointing
普通反向传播保存许多中间激活。Gradient Checkpointing 只保存部分节点,backward 时重新计算缺失激活:
更少激活内存
换取更多计算时间它主要节省激活,不减少权重本身。长序列和深层模型通常更受益。
8.6 数据并行、张量并行和流水线并行
- Data Parallel:每张卡有模型副本,处理不同 batch,聚合梯度。
- Tensor Parallel:一个大矩阵或一层拆到多张卡。
- Pipeline Parallel:不同层放到不同卡,micro-batch 像流水线经过各阶段。
- Sequence/Context Parallel:序列维度在设备间拆分。
- Expert Parallel:MoE 的专家分散到不同设备。
FSDP/ZeRO 的核心是把训练状态分片:
参数、梯度、优化器状态
不再每张卡完整复制
而是各设备只长期持有一部分代价是需要 all-gather、reduce-scatter 等通信。优化的核心变成计算、显存和网络通信的权衡。
8.7 Checkpoint 应包含什么
最低限度:
- 模型权重。
- 优化器状态。
- scheduler 状态。
- 当前 step/epoch。
- 随机数状态。
- 模型与训练配置。
- tokenizer。
- 数据游标或可恢复的数据顺序。
只保存模型权重可以用于推理,但不能保证精确恢复训练。
8.8 训练异常诊断
| 现象 | 常见原因 | 首先检查 |
|---|---|---|
| loss 不下降 | 标签错位、学习率不当、参数未更新 | batch、grad、optimizer |
| loss 立即很低 | 数据泄漏、未来 token 可见 | causal mask、labels |
| loss 变 NaN | 学习率过大、低精度溢出 | grad norm、激活范围 |
| 训练 loss 降,验证升 | 过拟合 | 数据划分、正则、训练轮数 |
| 速度忽快忽慢 | 数据加载、编译、swap | profiler、内存压力 |
| checkpoint 恢复后突变 | 状态缺失、数据顺序改变 | optimizer/scheduler/RNG |
8.9 规模估算
非常粗略的 dense Transformer 训练计算量估算:
训练 FLOPs ≈ 6 × 参数量 N × 训练 token 数 D该近似用于量级判断,不包含所有架构和实现细节。例如 1B 参数训练 100B token:
约 6 × 10^9 × 10^11 = 6 × 10^20 FLOPs这说明“模型权重放得下”与“能够合理时间训练完”完全不同。
8.10 动手设计:训练容量与恢复演练
设计实验|Training Runbook 不要求多卡;目标是把训练配置变成可计算、可恢复的工程计划。
选择一个假设模型,填写:
| 项目 | 数值 | 依据 |
|---|---|---|
| 参数量、层数、hidden、context | config 或模型卡 | |
| 权重 dtype | 训练配置 | |
| micro-batch、累积、设备数 | 有效 batch 公式 | |
| 每 step 有效 token | 非 padding token 实测 | |
| 目标训练 token | 数据和实验目标 | |
| 预计 optimizer step | 目标 token / 每 step token | |
| checkpoint 间隔 | 可接受的最大重算时间 | |
| 保留策略 | last、best、周期快照 |
恢复演练不需要等待真正故障:
- 固定 seed 训练 20 step,保存完整 checkpoint。
- 从 checkpoint 继续 10 step,记录 loss 与参数摘要。
- 从头运行相同 30 step,比较第 30 step 的结果。
- 再故意只加载模型权重,不加载 optimizer、scheduler 和 RNG,观察轨迹差异。
完全逐 bit 一致受设备和内核确定性影响;验收重点是状态是否完整、恢复后的学习率和数据位置是否正确,以及差异能否解释。
章节练习:
- micro-batch 不变、累积步数翻倍时,每 optimizer step 的 token 与更新频率怎样变化?
- 混合精度 OOM 时,为什么只切换 FP16/BF16 不一定解决激活内存?
- FSDP 减少每卡长期驻留状态时,新增了什么通信?
- loss NaN 后为什么不能直接跳过该 step 然后继续宣布训练正常?
验收标准:拿到任意训练配置,能估算 step、状态容量和 checkpoint 风险;面对异常能先检查数据/Mask、数值、优化器和系统资源,而不是立刻重启任务。
本章依据
原理性结论以原始论文、官方文档或公开教材为依据。论文中的实验结果只适用于其声明的模型、数据、硬件和评估设置。
计算最优预训练下参数规模和训练 token 的配比。
低精度计算、主权重与 loss scaling。
优化器状态、梯度和参数的分片策略。
激活重计算与 gradient checkpointing 的计算—内存交换。
Transformer 张量并行训练的代表性实现。