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

如何在BERT Trainer实例中添加early stopping功能?

给Hugging Face Trainer添加早停功能的实现方案

以下针对Transformers库的官方Trainer实例,提供两种可直接落地的实现思路:


方案1:使用官方内置EarlyStoppingCallback(最推荐,上手最快)

早停逻辑完全依赖验证集指标判断,因此首先要保证你的训练配置开启了验证流程:

  • 第一步导入依赖:
    from transformers import EarlyStoppingCallback
  • 第二步配置TrainingArguments时必须设置以下核心参数:
    • evaluation_strategy 和 save_strategy 保持一致,可选epoch或steps,控制验证和模型保存的频率
    • 必须开启load_best_model_at_end = True,否则早停触发后不会自动回退到表现最优的模型权重
    • metric_for_best_model 设为多分类任务要跟踪的指标,比如accuracy、f1_macro、eval_loss等
    • greater_is_better 根据跟踪指标设置,比如准确率、F1值越高越好就设为True,跟踪损失就设为False
  • 第三步初始化Trainer时传入回调即可,示例代码:
trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=train_dataset,
    eval_dataset=eval_dataset,
    compute_metrics=compute_metrics, # 你自定义的多分类指标计算函数
    # patience设为3意味着连续3个验证周期指标没有达标提升就停止训练
    callbacks=[EarlyStoppingCallback(early_stopping_patience=3)]
)
  • 可选参数补充:可以新增early_stopping_threshold参数设置指标的最小提升阈值,比如设为0.001就代表指标至少提升0.001才被判定为有效优化,避免微小波动干扰早停判断。

方案2:自定义早停回调(适合需要特殊判断逻辑的场景)

如果内置回调满足不了需求,比如要同时跟踪多个指标、或者加入自定义的停止规则,可以自己继承TrainerCallback实现逻辑:

  • 自定义回调示例代码:
from transformers import TrainerCallback

class CustomEarlyStoppingCallback(TrainerCallback):
    def __init__(self, patience=3, min_delta=0.001):
        self.patience = patience # 容忍的无提升验证轮数
        self.min_delta = min_delta # 最小提升阈值
        self.best_metric = None # 存储历史最优指标
        self.wait = 0 # 计数连续无提升的轮数

    def on_evaluate(self, args, state, control, metrics=None, **kwargs):
        # 替换为你要跟踪的多分类指标,比如eval_f1_macro等
        current_metric = metrics.get("eval_accuracy")
        if self.best_metric is None:
            self.best_metric = current_metric
        elif current_metric < self.best_metric + self.min_delta:
            self.wait += 1
            # 达到容忍阈值就触发停止
            if self.wait >= self.patience:
                control.should_training_stop = True
        else:
            # 指标更新,重置计数器
            self.best_metric = current_metric
            self.wait = 0
  • 后续将自定义回调传入Trainer的callbacks参数即可,用法和内置回调完全一致。

注意事项

  • 多分类场景下,你自定义的compute_metrics函数必须正确返回你要跟踪的指标键值对,否则早停回调找不到对应指标会直接报错
  • 必须保证evaluation_strategy和save_strategy的频率完全一致,否则load_best_model_at_end配置会触发报错
  • 分布式训练场景下,要保证所有进程的验证指标计算逻辑一致,避免早停判断出现偏差

内容的提问来源于stack exchange,提问作者soulwreckedyouth

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 19:45:02