Skip to content

第 7 章 用 Trainer 训练模型 ​

学习目标 ​

  • 掌握 TrainingArguments 的关键参数
  • 掌握 Trainer 的组装方式与 compute_metrics
  • 能逐行读懂训练日志
  • 会评估并保存训练好的模型

7.1 Trainer:高层训练 API ​

Trainer 是 transformers 提供的训练封装:你给出模型、数据集、超参数,它负责 batch 循环、梯度更新、日志、评估、检查点保存。对大多数微调任务,用它比自己手写循环更省事、更不容易出错。

组装一个 Trainer 需要四样东西:

python
from transformers import TrainingArguments, Trainer

args = TrainingArguments(output_dir="./results")

trainer = Trainer(
    model=model,                # 要训练的模型
    args=args,                  # 训练超参数
    train_dataset=train_set,    # 训练集
    eval_dataset=eval_set,      # 验证集
    data_collator=collator,          # 如何拼 batch
    compute_metrics=compute_metrics, # 如何算指标
)

7.2 TrainingArguments 关键参数 ​

TrainingArguments 集中管理所有训练超参数,常用:

参数含义本书示例
output_dir日志与检查点保存目录"./results"
num_train_epochs训练轮数2
per_device_train_batch_size每张卡每步的 batch 大小16
per_device_eval_batch_size评估时的 batch 大小32
learning_rate学习率1e-4
eval_strategy何时评估:"epoch" 或 "steps""epoch"
eval_stepseval_strategy="steps" 时的评估间隔50
logging_strategy / logging_steps日志间隔"steps" / 5
save_strategy检查点保存策略"no"
report_to日志平台;[] 表示不集成[]
fp16 / bf16混合精度fp16=True
disable_tqdm关闭进度条(干净输出)True

7.3 compute_metrics:接入评估指标 ​

compute_metrics 接收模型的预测结果,返回指标字典。用第 12 章会细讲的 evaluate 库计算准确率:

python
import numpy as np
from evaluate import load as load_metric

accuracy = load_metric("accuracy")

def compute_metrics(eval_pred):
    logits, labels = eval_pred
    preds = np.argmax(logits, axis=-1)
    return accuracy.compute(predictions=preds, references=labels)

小实验验证一下:

python
print(accuracy.compute(predictions=[0, 1, 1, 0],
                       references=[0, 1, 0, 0]))

输出:

{'accuracy': 0.75}

7.4 第一次训练:200 条样本的实验 ​

用 prajjwal1/bert-tiny(约 438 万参数)+ SST-2 前 200 条,2 个 epoch,完整代码:

python
from datasets import load_dataset
from transformers import (BertConfig, BertForSequenceClassification,
                          AutoTokenizer, TrainingArguments, Trainer,
                          DataCollatorWithPadding, set_seed)

set_seed(42)
tokenizer = AutoTokenizer.from_pretrained("google-bert/bert-base-uncased")

def tokenize(examples):
    return tokenizer(examples["sentence"], truncation=True, max_length=64)

def prepare(split):
    ds = load_dataset("nyu-mll/glue", "sst2", split=split)
    return (ds.map(tokenize, batched=True, remove_columns=["sentence", "idx"])
              .rename_column("label", "labels"))

train_set = prepare("train[:200]")
eval_set = prepare("validation[:100]")

config = BertConfig.from_pretrained("prajjwal1/bert-tiny", num_labels=2)
model = BertForSequenceClassification.from_pretrained(
    "prajjwal1/bert-tiny", config=config)

args = TrainingArguments(
    output_dir="./results-tiny",
    num_train_epochs=2,
    per_device_train_batch_size=16,
    per_device_eval_batch_size=32,
    eval_strategy="epoch",
    logging_strategy="steps",
    logging_steps=5,
    save_strategy="no",
    report_to=[],
    disable_tqdm=True,
    fp16=True,
)

collator = DataCollatorWithPadding(tokenizer=tokenizer, return_tensors="pt")
trainer = Trainer(model=model, args=args,
                  train_dataset=train_set, eval_dataset=eval_set,
                  data_collator=collator, compute_metrics=compute_metrics)

result = trainer.train()
print(result)

训练日志(真实运行,进度条已省略):

{'loss': '0.6956', 'grad_norm': '1.733', 'learning_rate': '4.231e-05', 'epoch': '0.3846'}
{'loss': '0.6936', 'grad_norm': '1.589', 'learning_rate': '3.269e-05', 'epoch': '0.7692'}
{'eval_loss': '0.6966', 'eval_accuracy': '0.48', 'eval_runtime': '0.0384', 'epoch': '1'}
{'loss': '0.688', 'grad_norm': '1.368', 'learning_rate': '2.308e-05', 'epoch': '1.154'}
{'loss': '0.6844', 'grad_norm': '1.589', 'learning_rate': '1.346e-05', 'epoch': '1.538'}
{'loss': '0.6823', 'grad_norm': '1.88', 'learning_rate': '3.846e-06', 'epoch': '1.923'}
{'eval_loss': '0.696', 'eval_accuracy': '0.48', 'eval_runtime': '0.0309', 'epoch': '2'}
TrainOutput(global_step=26, training_loss=0.688417429, metrics={...})

7.5 逐行读懂训练日志 ​

每行日志的含义:

  • loss:最近一个日志间隔的平均训练损失,应该整体下降。
  • grad_norm:梯度范数,过大说明梯度爆炸(可加梯度裁剪)。
  • learning_rate:当前学习率。默认调度器会随训练衰减。
  • epoch:当前进度,0.3846 = 完成了 38.46%。
  • eval_loss / eval_accuracy:每次评估的结果,eval_strategy="epoch" 时每轮评估一次。

在这个 200 条样本的实验里,eval_accuracy 只有 0.48(随机猜测是 0.5),损失几乎没有下降。这不是代码 bug,而是两个原因:样本太少、模型太小、训练不充分。下一章加大数据量后,效果立刻不同。

train() 返回 TrainOutput,常用字段:global_step(总步数,这里是 26)、training_loss。

7.6 评估与保存 ​

python
metrics = trainer.evaluate()
print(metrics)

输出:

{'eval_loss': 0.696, 'eval_accuracy': 0.48, 'eval_runtime': 0.0315,
 'eval_samples_per_second': 3178.7, 'eval_steps_per_second': 127.1, 'epoch': 2}

保存模型与分词器(第 3 章的 save_pretrained):

python
trainer.save_model("./my-finetuned-model")
tokenizer.save_pretrained("./my-finetuned-model")

Trainer.save_model 会保存 config、权重和训练参数,加上分词器文件后,这个目录就是一个完整的可加载模型。

动手实践 ​

  1. 把 7.4 的实验改成 train[:2000]、learning_rate=1e-4、1 个 epoch,观察 loss 与 eval_accuracy 的变化。
  2. 把 eval_strategy 从 "epoch" 改成 "steps"(配合 eval_steps=10),观察评估频率差异。
  3. 训练完成后用 trainer.evaluate() 输出指标,把结果与 7.4 对比。

常见错误 ​

错误 1:eval_strategy 与 eval_steps 不匹配。

eval_strategy="steps" 时必须给 eval_steps;"epoch" 时按轮评估,eval_steps 无效。

错误 2:忘记 report_to=[],本地没有 wandb 却报错。

集成平台没装会直接报 ImportError。教程示例统一用 report_to=[]。

错误 3:小数据集上迷信 eval_accuracy。

200 条样本、几百步训练出的数字没有统计意义(7.4 就是 0.48)。数字要在足够大的验证集 + 足够训练步数下才可信。

章末练习 ​

基础

  1. 列出 TrainingArguments 中控制「训练轮数」「batch 大小」「评估频率」的三个参数。
  2. 写一个 compute_metrics,返回 accuracy 与 f1 两个指标(提示:用 load_metric("f1"))。

提高

  1. 用 train[:200] 分别跑 num_train_epochs=1 和 3,对比 eval_accuracy,解释差异。
  2. 打开 output_dir 目录,找出训练产生的文件,说明 trainer_state.json 里存了什么。

挑战

  1. 给 Trainer 添加回调,在每个 epoch 结束时打印「当前 epoch + 当前 eval_accuracy」(提示:TrainerCallback 的 on_epoch_end)。

章末自测 ​

  1. Trainer 的 compute_metrics 参数接收什么、返回什么?
  2. eval_strategy="epoch" 表示什么?
  3. TrainOutput.global_step 表示什么?
  4. 判断:200 条样本上得到 0.48 的准确率说明代码写错了。
  5. 训练日志中 grad_norm 过大通常说明什么问题?
  6. report_to=[] 的作用是什么?