跳到主要内容

数据截至 (上游 commit 090253dac668)

03 · 训练器:万亿 token 长跑的中央循环

这一章讲什么: olmo/train.pyTrainer 类——一个训练 step 内部发生什么,以及更重要地,一个跑几个月的训练任务怎么存、怎么续、怎么停


1. 它要解决的小问题

「训练循环」在教程里是五行代码;在 4T token、几百张卡、跑几个月的现实里,真正难的是:

  • 显存装不下一个 batch → micro-batch 梯度累积,且 loss 归一化必须按 token 而非按 micro-batch 个数。
  • 梯度偶尔会炸 → 每步都要量梯度范数、按需裁剪,还要把这些数字全记下来。
  • 任务一定会挂(机器故障、抢占、loss 尖峰)→ checkpoint 要频繁、恢复要精确到「数据流的同一点 + 同一个 RNG 状态」。
  • 人要在不 SSH 的情况下停任务 → 最好有个网页按钮。

Trainerolmo/train.py:208)就是把这些全部收编的那个 dataclass。


2. 顶层结构:fit() 主循环

Trainer.fitolmo/train.py:1109)的骨架(省略监控与 profiler):

准备工作:
gc.disable() # FSDP 和自动 GC 处不好 (:1121-1123)
(可选) eval_on_load: 恢复后先评测一轮

for epoch in range(epoch, max_epochs):
for batch in train_loader:
① 记账与硬断言 global_step += 1; 断言 seq_len / batch_size (:1206-1208)
② train_step(batch) 前反向 + 裁剪 + 调度 + 优化器步进
③ 日志 console / W&B
④ 取消检测 每 50 步 check_if_cancelled() (:1256)
⑤ checkpoint 定期 sharded / ephemeral / unsharded (:1262-1309)
⑥ 评测 每 eval_interval 步 (:1312-1325)
⑦ 手动 gen1 GC 每步 gc.collect(1) (:1336-1337)
到 stop_at 就 break

收尾: 存最终 checkpoint (:1360-1376)

②里的硬断言值得单独说assert seq_len == cfg.model.max_sequence_lengthassert batch_size == cfg.device_train_batch_sizeolmo/train.py:1206-1208)。因为 token 计数(global_train_tokens_seen)是直接拿 batch 形状乘出来的,不做分布式 all-reduce;注释写明这是故意的——宁可违反假设时当场炸,也不要每次迭代付一次通信开销(:1199-1205)。


3. 核心机制一:一个 step 内部

train_stepolmo/train.py:832)的固定动作序列:

0. 写数据索引 TSV(:836-838) ← 审计线索,见第 2 章
1. zero_grad
2. train_batch(batch) ← 本节重点
3. dist.reduce(ce_loss) 到 rank 0 ← 只在「这步要记日志」时做 (:854)
4. clip_grads_and_collect_metrics() ← 量范数 + 裁剪,见 §4
5. scheduler.get_lr() 重算每组 LR ← 见 §5
6. optim.step()
7. NaN 检查 (:891-894) ← nan 直接 raise,长跑里宁可死也要死得早

3.1 micro-batch 与 token 归一化

train_batcholmo/train.py:782)先把 device batch 切成 device_train_microbatch_size 的小块(split_batch:934-954),逐块前反向。DDP 模式下除最后一块外包 no_sync:796-803),梯度只在最后同步一次——这是梯度累积的标准省通信手法。

loss 计算在 train_micro_batch:759-780),关键在这一行:

# 摘自 olmo/train.py:762-765
ce_loss, z_loss, logits = self.model_forward(
micro_batch, compute_z_loss=..., loss_reduction="sum"
)
ce_loss = ce_loss / batch_size_in_tokens

注意:loss 用 sum 规约,再除以整个 batch 的 token 数(不是 micro-batch 的,也不是实例数)。这样无论怎么切 micro-batch,梯度数学上都等于「整个 batch 上按 token 平均」——切分纯粹是显存技巧,不改训练语义。

3.2 z-loss 与 fused loss

cross_entropy_lossolmo/train.py:124-145)支持可选的 z-loss(PaLM 提出的辅助损失,惩罚 softmax 分母的 logsumexp 过大,压住 logit 漂移):z_squared = logits.logsumexp(-1).pow(2),乘 auxiliary_loss_multiplier 加进 loss。OLMo-2 开了它,系数 1e-5(configs/official-1124/OLMo2-7B-stage1.yaml:38-39)。

装了 flash-attn 时还能换 fused kernel 版交叉熵(fused_loss_fnolmo/train.py:156-202),省掉 logits 大张量的物化;代码里还处理了一个真实的依赖坑——flash-attn 2.5.8 把参数名从 ignored_index 改成了 ignore_index:164-170)。

3.3 label 的三重掩码

get_labelsolmo/train.py:715-729):label = input_ids 左移一位(GPT 标准做法),但支持三种掩码置 -100(不计 loss):

掩码来源语义
label_mask数据文件附带任意 token 级开关(SFT 只算 answer)
attention_maskpaddingpadding 不计 loss
instance_mask周期重复过滤整条实例不计 loss

4. 核心机制二:梯度裁剪,顺便把指标全收集了

小问题: FSDP 下梯度是分片的,怎么算全局范数?裁剪要不要逐参数做?这些数字本身也是训练健康的仪表盘。

Optimizer.clip_grads_and_collect_metricsolmo/optim.py:49)三步走:

  1. 本地量:逐参数算 min/max/sum/norm(梯度 + 参数 + 优化器状态)(:92-146)。
  2. 跨 rank 归约:范数平方 all_reduce 求和再开方得全局范数;min/max/sum 用普通 reduce:151-205)。注释里特别指出范数必须 all_reduce——因为每个 rank 都要拿到全局范数才能正确裁剪(:160-163)。
  3. 裁剪:两种模式(:226-242)。

两种裁剪模式的差别:

模式配置字段规则实现
固定全局裁剪max_grad_norm全局范数 > 阈值就整体缩放_do_global_fixed_clippingolmo/optim.py:322
自适应逐参数max_grad_norm_ratio每个参数裁到「自身历史范数指数平均 × ratio」_do_adaptive_clipping:254

自适应模式是长跑项目对付 loss 尖峰的工具:它对「这个参数平时的梯度多大」有记忆(grad_norm_exp_avg:289-311),突发尖峰会被单独按下去而误伤不到别的参数。OLMo-2 stage1 用的是固定裁剪 max_grad_norm: 1.0

一个性能细节:裁剪时无条件乘 clip_coef_clampedclamp 到 ≤1),而不是 if clip_coef < 1 分支——注释说明这是为了避免 host-device 同步olmo/optim.py:304-305:347-348)。


5. 核心机制三:LR 调度作为一等公民

调度器族在 olmo/optim.py:651 起,全部是「get_lr(initial_lr, step, max_steps) 一个函数」的朴素接口,但有两个项目级特色。

特色一:单位可选 step 或 token。 SchedulerUnits.tokensscheduler_current = global_train_tokens_seenolmo/train.py:304-320)。OLMo-2 stage1 的调度写的就是 token:t_warmup: 8388608000(约 84 亿 token warmup)、t_max: 5e12

特色二:为「中途重开」设计的 BoltOnWarmupScheduler。 stage2 从 stage1 的 checkpoint 继续训但要重置优化器,如果直接用原调度器,LR 会从一个高点直接往下走、没有 warmup。BoltOnWarmupScheduler.wrapolmo/optim.py:754-777)在任意已有调度器外面包一段「线性爬升段」,爬升的终点精确落在原调度曲线当前位置的值上(intercept),然后无缝接回原曲线。接线代码在 scripts/train.py:362-367

其余调度器一览(都有 warmup + alpha_f 终值比例):CosWithWarmupoptim.py:694)、LinearWithWarmup:713,stage2 用它配 alpha_f: 0 线性降到 0)、InvSqrtWithWarmup:732)、MaxScheduler(取两者较大,:743)、CosLinearEnvelope(annealing 用,:793)。


6. 核心机制四:checkpoint——保存容易,恢复才见功力

6.1 三种 checkpoint 的分工

CheckpointTypeolmo/config.py:892):

类型内容频率控制用途
sharded每 rank 一片模型+优化器save_interval + save_num_checkpoints_to_keep训练中途恢复,便宜
sharded_ephemeral同上,但最多留 1 个save_interval_ephemeral高频保险丝,防挂
unshardedrank 0 聚合成单文件save_interval_unsharded发布/转换/推理用

OLMo-2 7B 的组合是 250 步 ephemeral + 1000 步 unsharded,而永久 sharded 用 save_num_checkpoints_to_keep: 0 关掉了——发布的每千步 checkpoint 都来自 unsharded 这条线(configs/official-1124/OLMo2-7B-stage1.yaml:70-76;主循环里存 sharded 前有 save_num_checkpoints_to_keep != 0 的守卫,olmo/train.py:1263-1268)。「ephemeral 只留最新一个」的逻辑就在 _save_checkpoint 里:num_checkpoints_to_keep = 1olmo/train.py:463)。每次永久 checkpoint 落盘后还会把 ephemeral 全清掉(:1275-1277)。

6.2 保存的仪式

_save_checkpointolmo/train.py:446)的固定动作:先 zero_grad(避免把梯度也聚合进来,:468)→ 写 checkpoint → 维护 latest 符号链接(:495-506,处理了多节点共享 NFS 时两个 rank 同时建链接的竞态)→ 删超龄旧 checkpoint → barrier()

FullCheckpointer.save_checkpointolmo/checkpoint.py:620-679)里模型和优化器分开两次聚合写盘,注释说明是为了省 CPU 内存(:637)。

6.3 恢复的四件套

trainer_state_dictolmo/train.py:322-339)存的不只是 step 数:

epoch / global_step / 已见实例数 / 已见 token 数 ← 位置
checkpoints / unsharded / ephemeral 列表 ← 簿记
python / numpy / torch / cuda / mps 五套 RNG 状态 ← 随机性

load_trainer_state_dict:341-429)的恢复顺序:簿记 → 位置(必要时 dataset.reshuffle(epoch)、设 dataset.start_index)→ 把 LR 重置回配置值而非 checkpoint 里的值:408-417,这样可以用新超参续跑)→ 恢复 RNG(world size 变了就跳过并告警,:420-429)。

fast_forward_batchesolmo/config.py:999-1004)是 loss 尖峰后的标准操作:从上一个好 checkpoint 重启,但把数据流再往前跳过 N 个 batch——直接跳过当时可能有毒的那批数据

6.4 NaN 哨兵恢复法

FullCheckpointer.restore_checkpointolmo/checkpoint.py:681-739)在 FSDP 下恢复单文件 checkpoint 时不用框架的 load_state_dict,而是:

① 把所有 FlatParam 填成 NaN (:697-703)
② 逐参数把 checkpoint 里的对应切片 copy_ 进分片 (:714-729)
③ 断言没有任何 NaN 残留 (:731-739)

任何参数漏恢复都会当场报「contains NaNs, this is likely a bug」。把「静默带着随机权重继续训」这类最阴险的故障变成响亮crash。 优化器状态则按 local rank 轮流加载,每轮后 gc.collect() + empty_cache():753-759)——防止所有 rank 同时把优化器状态读进内存打爆机器。

6.5 第 0 步的防空洞

开训前 scripts/train.py:334-344 先存一个 pre-train checkpoint 再立刻读回来——磁盘满、权限错、checkpointer 配置错这类问题在第 0 步暴露,而不是第 5 万步。


7. 核心机制五:怎么停一个万卡任务

check_if_cancelledolmo/train.py:1057-1107)每 canceled_check_interval(默认 50)步检查三种停止信号:

信号机制来源
时间到time_limit_start_time 比较集群配额
早停cur_train_loss > early_stopping_factor * min_train_loss(warmup 后才生效)loss 失控自动止损
人工W&B run 被打上 cancel/canceled 标签(走 import/export API 查)网页操作

rank 0 检测、synchronize_flag 广播给所有 rank(:1093)。取消后不是立刻死extra_steps_after_cancel(默认 10)让任务存完 checkpoint 再多训几步——注释说明这是为了重启后指标曲线有重叠段(olmo/config.py:1260-1265)。细节控到这种程度,是真实长跑项目的样子。


8. 坑与边界

  • 恢复 RNG 要求 world size 不变olmo/train.py:420);换卡数续跑会告警跳过 RNG 恢复,数据顺序仍对(靠 start_index),但 dropout 等随机性不保证逐位一致。
  • token 计数假设 batch 等长(§2 的断言);drop_last: true 几乎是必须的,官方配置都这么写。
  • optim/ 前缀指标只进 W&B 不进控制台log_metrics_to_console 里过滤,olmo/train.py:964-989)——指标太多,控制台只留总梯度范数。
  • 评测会搞乱 torch.compile 的编译缓存,所以评测完 torch.compiler.reset():1051-1053)。
  • DDP 是二等公民:checkpoint 对 DDP 只存 unsharded(scripts/train.py:301-308),且 init_device 必须 cuda:158-159)。主力路径是 FSDP。
  • 每步一次 gc.collect(1) + 启动时 gc.disable():Python 自动 GC 和 FSDP 的通信节奏互相干扰,作者选择手动收一代垃圾(:1121-1123:1336-1337)。

9. 代码地图

主题文件路径符号名
中央循环olmo/train.py:1109Trainer.fit
单步olmo/train.py:832train_step
micro-batcholmo/train.py:782:934train_batchsplit_batch
lossolmo/train.py:124:759:715cross_entropy_losstrain_micro_batchget_labels
梯度裁剪与指标olmo/optim.py:49:254:322clip_grads_and_collect_metrics_do_adaptive_clipping_do_global_fixed_clipping
调度器olmo/optim.py:651:694:754SchedulerCosWithWarmupBoltOnWarmupScheduler
参数分组olmo/optim.py:829get_param_groups(decay / no-decay 分组)
checkpoint 保存olmo/train.py:446olmo/checkpoint.py:620_save_checkpointFullCheckpointer.save_checkpoint
checkpoint 恢复olmo/train.py:341olmo/checkpoint.py:681load_trainer_state_dictFullCheckpointer.restore_checkpoint
sharded checkpointer 族olmo/checkpoint.py:892:969:1474:1909TorchNewStyleShardedCheckpointerTorchLegacyShardedCheckpointerLocalShardedCheckpointerOlmoCoreCheckpointer
取消机制olmo/train.py:1057check_if_cancelled
启动编排scripts/train.py:54main(pre-train checkpoint 在 :334-344