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

为何无法同时使用EarlyStoppingCallback与load_best_model_at_end=False?五折训练如何存最优模型?

问题

我正在进行5折交叉训练,希望为Seq2SeqTrainer添加EarlyStoppingCallback,使训练在模型性能无提升时自动停止,但运行时出现错误:

AssertionError: EarlyStoppingCallback requires load_best_model_at_end = True

我使用的代码如下:

training_args = Seq2SeqTrainingArguments(
    output_dir="./logs",
    evaluation_strategy="epoch",
    logging_strategy="epoch",
    learning_rate=2e-5,
    per_device_train_batch_size=16,
    per_device_eval_batch_size=16,
    weight_decay=0.01,
    num_train_epochs=5,
    save_total_limit=2,
    save_strategy="epoch",
    load_best_model_at_end=True,
    predict_with_generate=True,
    fp16=False,
    push_to_hub=False,
)

for train_dataset, val_dataset in zip(train_ds, val_ds):
    trainer = Seq2SeqTrainer(
        model=model,
        args=training_args,
        train_dataset=train_dataset,
        eval_dataset=val_dataset,
        tokenizer=tokenizer,
        data_collator=data_collator,
        compute_metrics=compute_metrics,
        callbacks=[
            CombinedTensorBoardCallback,
            EarlyStoppingCallback(early_stopping_patience=3),
        ],
    )

    train_result = trainer.train()

请问为何无法同时使用EarlyStoppingCallback并设置load_best_model_at_end=False?我仅想在每个折的训练阶段保存最佳模型,另外,是否有方法在5折训练后保存所有折中表现最优的模型?


解答

为什么不能同时使用EarlyStoppingCallback和load_best_model_at_end=False

EarlyStoppingCallback的设计逻辑和load_best_model_at_end强绑定:

  • 这个回调的核心是追踪验证集的性能变化,当连续early_stopping_patience个epoch没有提升时终止训练。
  • 要实现这个逻辑,训练器需要持续记录训练过程中性能最优的模型权重——而load_best_model_at_end=True正是开启这个追踪机制的开关。如果设为False,训练器不会保存和追踪最佳模型状态,回调也就无法保证停止时模型处于最优状态,因此会触发断言错误。

其实你想保存每个折的最佳模型,开启load_best_model_at_end=True反而能满足需求:训练结束后,Trainer会自动将模型加载为该折的最佳权重,此时调用trainer.save_model()就能直接保存该折的最优模型。

5折训练后保存所有折中表现最优的模型

可以通过以下步骤实现:

  1. 为每个折分配独立的输出目录,避免模型文件互相覆盖。
  2. 每个折训练完成后,记录该折的最佳验证指标,同时保存该折的最佳模型。
  3. 所有折训练结束后,对比所有折的验证指标,选出最优的那个模型作为最终结果保存。

修改后的代码示例:

import shutil

# 初始化列表记录各折的结果
fold_results = []
# 遍历每个折
for fold_idx, (train_dataset, val_dataset) in enumerate(zip(train_ds, val_ds)):
    # 为当前折设置独立的输出目录
    fold_output_dir = f"./logs/fold_{fold_idx}"
    training_args.output_dir = fold_output_dir
    
    trainer = Seq2SeqTrainer(
        model=model,
        args=training_args,
        train_dataset=train_dataset,
        eval_dataset=val_dataset,
        tokenizer=tokenizer,
        data_collator=data_collator,
        compute_metrics=compute_metrics,
        callbacks=[
            CombinedTensorBoardCallback(),  # 注意:这里需要实例化,原代码可能存在问题
            EarlyStoppingCallback(early_stopping_patience=3),
        ],
    )

    train_result = trainer.train()
    
    # 获取当前折的最佳验证指标(根据你的compute_metrics返回值调整,比如bleu、loss等)
    best_metric = trainer.state.best_metric
    # 设置当前折最佳模型的保存路径
    best_model_path = f"./best_fold_models/fold_{fold_idx}"
    # 保存当前折的最佳模型
    trainer.save_model(best_model_path)
    
    # 将当前折的结果存入列表
    fold_results.append({
        "fold_idx": fold_idx,
        "best_metric": best_metric,
        "model_path": best_model_path
    })

# 筛选出最优模型(如果是loss则用min,这里假设指标越大越好,比如bleu)
best_fold = max(fold_results, key=lambda x: x["best_metric"])
# 将最优模型复制到最终保存目录
shutil.copytree(best_fold["model_path"], "./final_best_model", dirs_exist_ok=True)
print(f"最优模型来自第{best_fold['fold_idx']}折,最佳指标:{best_fold['best_metric']},已保存到./final_best_model")

注意:原代码中的CombinedTensorBoardCallback需要实例化(加括号),否则会报错,上面的示例已修正。


内容的提问来源于stack exchange,提问作者Houcem Ben Makhlouf

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 19:40:06