You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.15 23:43:15