如何在Hugging Face Transformers Trainer中获取每个epoch或step的训练准确率?
嗨,这个问题我之前做序列分类任务时也碰到过——Hugging Face的Trainer默认确实只会在训练日志里输出损失值,不会主动计算并记录训练准确率。不过有几种实用的方法可以解决这个需求,我给你详细说说:
方法一:重写Trainer的训练步骤,实时记录批次准确率
这种方法直接继承Trainer类,重写training_step方法,在每一步训练时计算当前批次的准确率,并把它和损失一起写入日志。
示例代码:
from transformers import Trainer import torch class CustomTrainer(Trainer): def training_step(self, model, inputs): model.train() inputs = self._prepare_inputs(inputs) # 前向传播计算损失和模型输出 with self.compute_loss_context_manager(): loss, outputs = self.compute_loss(model, inputs, return_outputs=True) # 计算当前训练批次的准确率 logits = outputs.logits preds = torch.argmax(logits, dim=-1) labels = inputs["labels"] train_accuracy = (preds == labels).float().mean().item() # 将准确率和损失一起记录到日志 self.log({"train_loss": loss.item(), "train_accuracy": train_accuracy}) # 保留原Trainer的梯度更新逻辑 loss = loss / self.args.gradient_accumulation_steps if self.args.use_apex: with amp.scale_loss(loss, self.optimizer) as scaled_loss: scaled_loss.backward() else: loss.backward() return loss.detach() # 替换原Trainer为自定义的CustomTrainer trainer = CustomTrainer( model=model, args=training_args, train_dataset=train_dataset, eval_dataset=val_dataset, compute_metrics=compute_metrics )
修改后,每一步训练都会输出当前批次的train_accuracy,你可以在控制台或者logging_dir下的日志文件里看到这个指标。
方法二:用TrainerCallback实现灵活的准确率记录
如果你不想重写整个Trainer类,可以用Hugging Face的回调机制,在训练步骤结束后触发准确率计算和日志记录。这种方法更灵活,还能自定义记录频率(比如每隔10步记录一次)。
示例代码:
from transformers import TrainerCallback, TrainerState, TrainerControl class TrainAccuracyCallback(TrainerCallback): def on_step_end(self, args, state: TrainerState, control: TrainerControl, **kwargs): # 可以自定义记录频率,比如和logging_steps保持一致 if state.global_step % args.logging_steps != 0: return control # 获取模型和当前批次输入 model = kwargs["model"] inputs = kwargs["inputs"] # 用eval模式计算准确率,避免梯度更新 model.eval() with torch.no_grad(): outputs = model(**inputs) # 计算准确率 logits = outputs.logits preds = torch.argmax(logits, dim=-1) labels = inputs["labels"] train_accuracy = (preds == labels).float().mean().item() # 记录到日志 kwargs["trainer"].log({"train_accuracy": train_accuracy}) return control # 将回调添加到Trainer中 trainer = Trainer( model=model, args=training_args, train_dataset=train_dataset, eval_dataset=val_dataset, compute_metrics=compute_metrics, callbacks=[TrainAccuracyCallback()] )
这个回调会在每logging_steps步结束后计算并记录训练准确率,既满足你的日志需求,又不会额外增加太多计算开销。
一些注意点
- 确保你的训练数据集的
inputs里包含labels字段,这是计算准确率的必要条件,序列分类任务的数据集一般都会处理好这个,但如果是自定义数据集可以检查一下。 - 计算训练准确率会增加一定的计算量,如果你的数据集很大或者硬件资源有限,建议用方法二的频率控制,避免每步都计算。
- 如果你想记录整个epoch的训练准确率,可以把回调的触发时机改成
on_epoch_end,然后在该方法里遍历整个训练集计算准确率(不过这种方式计算量较大,谨慎使用)。
内容的提问来源于stack exchange,提问作者CptBaas
相关产品推荐
相关产品推荐

