GOCLAWLLM ENGINEERING
GoClaw 首页

8. 预训练工程

工程进阶2~3 小时
学习目标
  1. 拆解一个训练 step 的计算与内存
  2. 区分 batch、累积和并行策略
  3. 诊断 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

训练通常比推理更吃内存,因为不仅需要权重,还需要:

FP16 权重约 2 字节/参数,但使用 AdamW 的完整训练总占用可能远高于 2 字节/参数。

8.3 Batch、Micro-batch 和梯度累积

显存放不下目标 batch 时:

有效 batch
= 每设备 micro-batch
  × gradient_accumulation_steps
  × 数据并行设备数

梯度累积在多次 micro-batch 后才执行 optimizer step。若 loss 默认取均值,框架通常会正确缩放;手写训练时要注意除以累积步数。

梯度累积能减少峰值内存,但不能完全等价于大 batch:

8.4 混合精度

混合精度并不是把所有东西都盲目转成低精度。常见做法是矩阵乘法使用低精度,某些累加、归一化或优化器状态保留更高精度。

Apple MPS 与 NVIDIA CUDA 支持细节不同。不要把 CUDA 教程中的 bitsandbytestorch.cuda.amp 参数原样套到 Mac。

8.5 Gradient Checkpointing

普通反向传播保存许多中间激活。Gradient Checkpointing 只保存部分节点,backward 时重新计算缺失激活:

更少激活内存
换取更多计算时间

它主要节省激活,不减少权重本身。长序列和深层模型通常更受益。

8.6 数据并行、张量并行和流水线并行

FSDP/ZeRO 的核心是把训练状态分片:

参数、梯度、优化器状态
不再每张卡完整复制
而是各设备只长期持有一部分

代价是需要 all-gather、reduce-scatter 等通信。优化的核心变成计算、显存和网络通信的权衡。

8.7 Checkpoint 应包含什么

最低限度:

只保存模型权重可以用于推理,但不能保证精确恢复训练。

8.8 训练异常诊断

现象常见原因首先检查
loss 不下降标签错位、学习率不当、参数未更新batch、grad、optimizer
loss 立即很低数据泄漏、未来 token 可见causal mask、labels
loss 变 NaN学习率过大、低精度溢出grad norm、激活范围
训练 loss 降,验证升过拟合数据划分、正则、训练轮数
速度忽快忽慢数据加载、编译、swapprofiler、内存压力
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、contextconfig 或模型卡
权重 dtype训练配置
micro-batch、累积、设备数有效 batch 公式
每 step 有效 token非 padding token 实测
目标训练 token数据和实验目标
预计 optimizer step目标 token / 每 step token
checkpoint 间隔可接受的最大重算时间
保留策略last、best、周期快照

恢复演练不需要等待真正故障:

  1. 固定 seed 训练 20 step,保存完整 checkpoint。
  2. 从 checkpoint 继续 10 step,记录 loss 与参数摘要。
  3. 从头运行相同 30 step,比较第 30 step 的结果。
  4. 再故意只加载模型权重,不加载 optimizer、scheduler 和 RNG,观察轨迹差异。

完全逐 bit 一致受设备和内核确定性影响;验收重点是状态是否完整、恢复后的学习率和数据位置是否正确,以及差异能否解释。

章节练习:

验收标准:拿到任意训练配置,能估算 step、状态容量和 checkpoint 风险;面对异常能先检查数据/Mask、数值、优化器和系统资源,而不是立刻重启任务。


本章依据

原理性结论以原始论文、官方文档或公开教材为依据。论文中的实验结果只适用于其声明的模型、数据、硬件和评估设置。

  1. 计算最优预训练下参数规模和训练 token 的配比。

  2. 低精度计算、主权重与 loss scaling。

  3. 优化器状态、梯度和参数的分片策略。

  4. 激活重计算与 gradient checkpointing 的计算—内存交换。

  5. Transformer 张量并行训练的代表性实现。