跳到主要内容

数据截至 (上游 commit b6c0bfe04c82)

05 · Trainer 训练回路

这一章讲什么: 全生态微调代码坐的那层底座。读完你会知道 Trainer 到底替你做掉了什么、compute_loss 该在哪重写、以及梯度累积下 loss 为什么曾经算错、现在怎么修。


1. 它要解决的小问题

手写一个训练循环只要十行。但生产训练循环还要处理:分布式启动与通信、混合精度、梯度累积与裁剪、checkpoint 存/续训(含 dataloader 快进)、日志与评测穿插、早停、几十种优化器和调度器、以及「换台机器别改代码」。

Trainer(trainer.py:258)的答案是:把「流程编排」和「分布式执行」分开——后者整包委托给 accelerateAccelerator(在 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()

两个扩展点:

  1. compute_loss(:1980)是默认的重写点。 默认实现:若模型吃 loss kwargs 就把 num_items_in_batch 塞进 forward(:2018-2020),outputs = model(**inputs) 后取 outputs.loss——也就是说默认 loss 是模型自己算的(回想第 1 章 LlamaForCausalLM.forwardself.loss_function(...))。TRL 的 DPO/SFT 就是在这里换 loss。
  2. 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 over labels[..., 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)里与分布式相关的三个动作:

  1. 每个累积窗内,前 N-1 个 micro-batch 用 accelerator.no_sync(model) 包住(:1772 一带)——DDP 下跳过梯度 all-reduce,只在最后一个 micro-batch 同步,省掉 (N-1)/N 的通信。
  2. do_sync_step = (step + 1) % gradient_accumulation_steps == 0 or (step + 1) == steps_in_epoch(:1757)——判断本 micro-batch 是否收尾;因为 accelerate 有梯度预取,还要手动 _set_sync_gradients(:1759)。
  3. 同步步的固定动作: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_trainedsteps_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.pyTrainer.train_inner_training_loop_run_epoch
一个训练步src/transformers/trainer.pyTrainer.training_stepcompute_loss_prepare_inputs
batch 收集与 token 计数src/transformers/trainer.pyget_batch_samples_get_num_items_in_batch
优化器/调度器创建src/transformers/trainer.pycreate_optimizer_and_schedulercreate_optimizer
accelerate 委托点src/transformers/trainer.pycreate_accelerator_and_postprocessself.accelerator
超参数全集src/transformers/training_args.pyTrainingArguments
回调体系src/transformers/trainer_callback.pyTrainerCallbackCallbackHandlerTrainerStateTrainerControl
seq2seq 变体src/transformers/trainer_seq2seq.pySeq2SeqTrainer