如何在Hugging Face Trainer/SFT Trainer中记录Step 0训练损失
关于Hugging Face Trainer/SFTTrainer记录Step 0训练损失的问题
核心结论
目前Trainer和SFTTrainer没有内置的类似eval_on_start的训练初始损失记录功能,自定义回调或手动记录是当前的可行方案,其中优化后的自定义回调更贴合Trainer的原生日志逻辑。
现有方案分析与优化
方案1:自定义回调(优化版)
你之前的回调实现可以进一步优化,无需硬编码第三方日志工具(如WandB),而是利用Trainer原生的日志系统,同时可以选择计算真实的初始损失值:
from transformers import TrainerCallback import torch class TrainOnStartCallback(TrainerCallback): def on_train_begin(self, args, state, control, logs=None, **kwargs): trainer = kwargs.get("trainer") if not trainer: return # 可选:计算真实的初始训练损失(替换None为实际值) initial_loss = self._calculate_initial_loss(trainer) # 构造Step 0的日志数据 logs = logs or {} logs["train/loss"] = initial_loss logs["train/global_step"] = 0 # 使用Trainer原生日志方法,自动适配所有日志后端(WandB、TensorBoard等) trainer.log(logs) def _calculate_initial_loss(self, trainer): # 从训练数据加载一个batch并计算损失 train_dataloader = trainer.get_train_dataloader() batch = next(iter(train_dataloader)) batch = trainer._prepare_inputs(batch) with torch.no_grad(): outputs = trainer.model(**batch) loss = outputs.loss.item() return loss
优化点说明:
- 借助
trainer.log()方法,无需手动调用特定日志工具,兼容Trainer支持的所有日志系统 - 新增
_calculate_initial_loss方法,可以获取真实的初始模型损失(而非占位的None) - 利用Trainer内置的
get_train_dataloader()和_prepare_inputs()方法,避免重复实现数据加载和设备迁移逻辑
方案2:手动记录
如果仅需要记录占位式的Step 0损失,手动记录是最简洁的方式:
# 若需真实损失,可提前计算后替换None trainer.log({"train/loss": None, "train/global_step": 0}) trainer.train()
注:这里建议使用trainer.log()而非直接调用WandB,保持日志系统的一致性
补充说明
目前Hugging Face官方并未提供类似eval_on_start的train_on_start参数,因此上述两种方案是当前的最佳实践:
- 若只需占位记录:手动记录最简便
- 若需真实初始损失或统一日志流程:优化后的自定义回调更合适
内容的提问来源于stack exchange,提问作者Charlie Parker
相关产品推荐
相关产品推荐

