Skip to content

训练主流程

训练不是一个动作,而是一组固定步骤反复执行。

一次 batch 的训练

最核心的训练代码可以抽象成:

python
outputs = model(**batch)
loss = outputs.loss

loss.backward()
optimizer.step()
optimizer.zero_grad()

含义是:

代码含义
model(**batch)前向传播,得到 logitsloss
loss.backward()根据 loss 反向计算梯度
optimizer.step()根据梯度更新模型参数
optimizer.zero_grad()清空旧梯度,准备下一批

这就是训练的心脏。

epoch、batch、step

三个词要分清楚:

名词含义
batch一小批样本,比如 8 条、32 条、128 条
step通常指一次参数更新,也就是训练一个 batch
epoch完整看完一遍训练集

如果训练集有 1000 条,batch_size = 50

text
1 epoch = 1000 / 50 = 20 steps

如果训练 2 个 epoch:

text
总 step 数约为 20 * 2 = 40

Trainer.train() 在做什么

项目里的 Trainer.train() 管的是全局节奏:

python
self._load_checkpoint()
train_dataloader = self._create_dataloader(...)

for epoch in range(self.config.epochs):
    self.model.train()

    for batch in train_dataloader:
        loss = self._train_one_step(batch)

        if self.global_step % self.config.save_steps == 0:
            metrics = self.evaluate()
            self._check_early_stop(metrics)
            self._save_checkpoint()

        self.global_step += 1

它不是只训练,还夹着:

  • checkpoint 恢复
  • dataloader 创建
  • 每个 batch 的训练
  • 定期验证
  • 保存 checkpoint
  • early stop 判断

save_steps 是什么

save_steps 的意思是:

text
每训练多少个 step,做一次 evaluate + save

它不是 epoch。

如果你希望每个 epoch 验证一次,可以粗略设成:

python
save_steps = len(train_dataloader)

比如训练集 1000 条,batch_size = 50

text
len(train_dataloader) = 20
save_steps = 20

这就接近每个 epoch 验证一次。

训练和验证不是同时发生

在训练循环里,验证是“插入”进去的:

text
训练一段 batch
  -> 暂停训练
  -> 跑完整个验证集
  -> 统计 loss / acc / f1
  -> 切回训练模式
  -> 继续训练

验证阶段不会更新参数。它只是考试。