数据截至 (上游 commit b6c0bfe04c82)
05 · Trainer 训练回路
这一章讲什么: 全生态微调代码坐的那层底座。读完你会知道 Trainer 到底替你做掉了什么、
compute_loss该在哪重写、以及梯度累积下 loss 为什么曾经算错、现在怎么修。
1. 它要解决的小问题
手写一个训练循 环只要十行。但生产训练循环还要处理:分布式启动与通信、混合精度、梯度累积与裁剪、checkpoint 存/续训(含 dataloader 快进)、日志与评测穿插、早停、几十种优化器和调度器、以及「换台机器别改代码」。
Trainer(trainer.py:258)的答案是:把「流程编排」和「分布式执行」分开——后者整包委托给 accelerate 的 Accelerator(在 create_accelerator_and_postprocess 里创建,trainer.py:771;实例化在 :829),Trainer 自己只剩编排和事件。
2. 顶层结构:三层循环 + 一条事件链
trainer.train() trainer.py:1350
└─ _inner_training_loop() trainer.py:1465
初始化:算步数 → 建/复 TrainerState → 建 optimizer/scheduler
→ accelerator.prepare 包模型(分布式/混精度在这发生)
│
└─ for epoch in range(epochs_trained, num_train_epochs) :1538
└─ _run_epoch() :1693
for step: get_batch_samples() :2139
└─ training_step() :1907
compute_loss → accelerator.backward :1980/:1976
└─ 攒够累积步数:optimizer.step → scheduler.step
→ global_step += 1 → on_step_end :1800-1814
└─ _maybe_log_save_evaluate(...) :1815
读法: 外层两层循环几乎不含逻辑,所有「做什么」都收敛到 training_step 和回调事件;所有「怎么做」(多卡同步、混精度 backward)都是 self.accelerator.*。
回调链由 CallbackHandler(trainer_callback.py:429)驱动,事件包括 on_epoch_begin/on_step_begin/on_optimizer_step/on_step_end/on_log/on_save 等。回调之间通过两个共享对象通信:TrainerState(:35,只读事实:global_step、日志历史)和 TrainerControl(:234,可写意图:should_training_stop 等)——回调把 control.should_training_stop = True 一置,主循环在下个检查点退出(:1828-1830)。
3. 一个训练步:training_step 解剖
training_step(trainer.py:1907)的骨架(示意其顺序):
# 示意,非源码
def training_step(self, model, inputs, num_items_in_batch=None):
inputs = self._prepare_inputs(inputs) # 搬到设备、类型对齐
with self.compute_loss_context_manager(): # 混精度 autocast 等
loss = self.compute_loss(model, inputs, num_items_in_batch=num_items_in_batch)
if 模型不认 num_items_in_batch:
loss = loss / self.current_gradient_accumulation_steps # 老式归一化兜底
self.accelerator.backward(loss) # 分布式/混精度 backward
return loss.detach()
两个扩展点:
compute_loss(:1980)是默认的重写点。 默认实现:若模型吃 loss kwargs 就把num_items_in_batch塞进 forward(:2018-2020),outputs = model(**inputs)后取outputs.loss——也就是说默认 loss 是模型自己算的(回想第 1 章LlamaForCausalLM.forward里self.loss_function(...))。TRL 的 DPO/SFT 就是在这里换 loss。compute_loss_func参数:不想动模型时,传一个外部 loss 函数,Trainer 会把 labels pop 出来交给它(:2012-2017)。
4. num_items_in_batch:梯度累积的 loss 修正
4.1 小问题
causal LM 的 loss 是「对一个 batch 里所有非 pad token 求平均」。开梯度累积后,每个 micro-batch 各自平均再相加——token 多的 micro-batch 和 token 少的 micro-batch 被等权了,正确做法是先求全部 token 的总 loss,再除以整个累积窗的 token 总数。
4.2 实现
取 batch 时就先把数算好。get_batch_samples(:2139)除了从 dataloader 攒出 gradient_accumulation_steps 个 micro-batch,还调 _get_num_items_in_batch(:2156):
- 统计口径是
labels.ne(-100).sum()——-100是「不参与 loss」的约定值(:2193)。 - causal LM 要 shift labels,所以数的时候用
labels[..., 1:]或 collator 给的shift_labels,避免多数一个位置(:2183-2192,注释明说 "count overlabels[..., 1:]to avoid over-counting position 0")。 - 多卡时
accelerator.gather(...).sum()合成全局 token 数(:2197-2199)。
这个数被一路传进 compute_loss → 模型 forward → self.loss_function(..., num_items_in_batch=...),loss 直接按全局 token 数归一。只有当模型不认这个 kwarg 时才退回老路「除以累积步数」(training_step 里 :1967-1970 的兜底,注释自称 "GA loss bug is not fixed during compute loss")。
为什么值得学: 这是一个「统计量要沿着数据管道提前算好、穿过三层调用送到位」的典型修法——不动循环结构,只加一个可空的随行参数。
5. 梯度累积与同步的编排
_run_epoch(:1693)里与分布式相关的三个动作:
- 每个累积窗内,前 N-1 个 micro-batch 用
accelerator.no_sync(model)包住(:1772 一带)——DDP 下跳过梯度 all-reduce,只在最后一个 micro-batch 同步,省掉 (N-1)/N 的通信。 do_sync_step = (step + 1) % gradient_accumulation_steps == 0 or (step + 1) == steps_in_epoch(:1757)——判断本 micro-batch 是否收尾;因为 accelerate 有梯度预取,还要手动_set_sync_gradients(:1759)。- 同步步的固定动作:
optimizer.step()(:1800)→on_optimizer_step回调(:1801)→ 没被跳过(梯度溢出等)就lr_scheduler.step()→global_step += 1(:1812)→on_step_end(:1814)→_maybe_log_save_evaluate(:1815)按步数决定要不要打日志/评测/存 checkpoint。
还有一个细节:累积窗最后一个窗若不满,current_gradient_accumulation_steps 被设为实际 micro-batch 数(:1743-1745),保证 §4 的归一化兜底仍然正确。
6. 断点续训
_init_training_state(trainer.py:1558 起)从 checkpoint 目录读 trainer_state.json 恢复 TrainerState,然后算 epochs_trained 和 steps_trained_in_current_epoch;默认还会快进 dataloader 到对应位置(_inner_training_loop 开头的日志,:1516-1521,"Fast-forwarding the dataloader past N epochs and M batches")。
注意这个快进是「重新消费并丢弃」已训过的 batch——对 map-style dataset 没问题,对 IterableDataset 要靠 ignore_data_skip 或数据集自身的状态恢复。
7. 关键细节与坑
- 重写点优先级:
compute_loss>training_step> 整个train。 改 loss 重写compute_loss即可;要动 backward/多模型才碰training_step;train几乎不用碰。 num_items_in_batch只对「模型认 loss kwargs」生效。 自定义模型若 forward 不声明**loss_kwargs,Trainer 会退回避坑归一化——发现 loss 曲线和理论对不上时先查model_accepts_loss_kwargs。- 回调里想停训练,置
control不要 raise。TrainerControl就是为此设计的;异常会跳过收尾逻辑(最终 checkpoint、_finalize_training)。 - Trainer 不管 3D 并行的细活。 TP/PP 的手工切分、流水调度不在它的抽象里;DeepSpeed/FSDP 是通过
TrainingArguments透传给 accelerate/DeepSpeed 的,行为细节要看那两边的文档。 auto_find_batch_size是 OOM 重启循环,不是免费午餐:每次 OOM 砍半 batch 重头来,已训的步数作废。
8. 代码地图
| 主题 | 文件路径 | 符号名 |
|---|---|---|
| 训练入口与主循环 | src/transformers/trainer.py | Trainer.train、_inner_training_loop、_run_epoch |
| 一个训练步 | src/transformers/trainer.py | Trainer.training_step、compute_loss、_prepare_inputs |
| batch 收集与 token 计数 | src/transformers/trainer.py | get_batch_samples、_get_num_items_in_batch |
| 优化器/调度器创建 | src/transformers/trainer.py | create_optimizer_and_scheduler、create_optimizer |
| accelerate 委托点 | src/transformers/trainer.py | create_accelerator_and_postprocess、self.accelerator |
| 超参数全集 | src/transformers/training_args.py | TrainingArguments |
| 回调体系 | src/transformers/trainer_callback.py | TrainerCallback、CallbackHandler、TrainerState、TrainerControl |
| seq2seq 变体 | src/transformers/trainer_seq2seq.py | Seq2SeqTrainer |