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

MarianMT微调添加EarlyStopping后Seq2SeqTrainer无报错崩溃求助

问题描述

使用Hugging Face的Seq2SeqTrainer对MarianMTModel进行翻译任务微调时,原本训练正常,但添加EarlyStoppingCallback后,进程无任何错误信息或回溯日志直接静默崩溃。

环境配置

  • 模型:MarianMTModel
  • 分词器:MarianTokenizer
  • 训练器:Seq2SeqTrainer
  • 评估指标:BLEU、METEOR(用于早停判断)
  • 批次大小:1(已因内存问题调至最小)

此前已从默认Trainer切换为Seq2SeqTrainer优化内存,但添加早停后仍出现崩溃。核心代码如下:

class MarianFineTuner:
    def __init__(self, model_name: str, device: str, config: dict):
        self.model_name = model_name
        self.device = device
        self.config = config
        self.tokenizer = MarianTokenizer.from_pretrained(model_name)
        self.model = MarianMTModel.from_pretrained(model_name).to(device)
        self.data_collator = DataCollatorForSeq2Seq(tokenizer=self.tokenizer, model=self.model)

    def tokenize_dataset(self, dataset, source_col: str, target_col: str):
        def tokenize_function(examples):
            model_inputs = self.tokenizer(examples[source_col], truncation=True)
            with self.tokenizer.as_target_tokenizer():
                labels = self.tokenizer(examples[target_col], truncation=True)
            model_inputs["labels"] = labels["input_ids"]
            return model_inputs

        return dataset.map(tokenize_function, batched=True, remove_columns=[source_col, target_col])


    def train(self, train_dataset, val_dataset, experiment_name):

        training_args = Seq2SeqTrainingArguments(
            output_dir=self.config["temp_output_dir"],
            per_device_train_batch_size=self.config["batch_size"],
            per_device_eval_batch_size=self.config["batch_size"],
            num_train_epochs=self.config["num_train_epochs"],
            eval_strategy="epoch",
            logging_strategy="epoch",
            save_strategy="epoch",
            save_total_limit=1,
            do_train=True,
            do_eval=True,
            report_to=[],
            load_best_model_at_end=True,
            metric_for_best_model="meteor",
            greater_is_better=True,
            predict_with_generate=True,
            torch_empty_cache_steps=2,
            eval_accumulation_steps=10,
        )

        log_path = f"results/epoch_logs/epoch_log_{experiment_name}.csv"

        trainer = Seq2SeqTrainer(
            model=self.model,
            args=training_args,
            train_dataset=train_dataset,
            eval_dataset=val_dataset,
            tokenizer=self.tokenizer,
            data_collator=self.data_collator,
            callbacks=[
                CSVLoggerCallback(log_path, experiment_name),
                EarlyStoppingCallback(early_stopping_patience=2)
            ],
            compute_metrics=self.compute_metrics
        )

        trainer.train()

分析与解决建议

1. 内存溢出(最可能原因)

虽然批次大小设为1,但开启predict_with_generate=True后,模型评估阶段会执行生成式推理(默认用beam search,beam size=5),这会占用远多于普通前向传播的内存:

  • 生成过程中需要维护多个候选序列的缓存
  • MarianMT的decoder会重复计算多个token的概率
  • 系统内存/显存不足时,会直接终止进程而不输出报错(这是Linux/Windows OOM killer的典型行为)

解决方法:

  • 降低beam search大小:在训练参数里添加generation_config=GenerationConfig(beam_size=1)(改用贪心搜索,内存占用大幅降低)
  • 启用混合精度训练:添加fp16=True到Seq2SeqTrainingArguments(需支持CUDA的GPU)
  • 开启梯度检查点:初始化模型后调用model.gradient_checkpointing_enable(),减少显存占用
  • 手动清理缓存:在compute_metrics函数开头添加torch.cuda.empty_cache()

2. 指标计算与早停逻辑不兼容

设置metric_for_best_model="meteor"后,若compute_metrics函数返回的字典中没有"meteor"键,或计算METEOR时出现未捕获的异常,可能触发早停逻辑崩溃:

  • 检查compute_metrics是否正确返回包含"meteor"的结果字典
  • 确保METEOR依赖库(如sacrebleu或nltk)已正确安装且版本兼容

3. 回调函数冲突或配置问题

自定义CSVLoggerCallback与EarlyStoppingCallback可能存在执行顺序或资源竞争问题:

  • 暂时移除CSVLoggerCallback,单独测试早停功能是否正常
  • 确认EarlyStoppingCallback参数合法(如early_stopping_patience=2是有效值)

4. 训练参数矛盾设置

load_best_model_at_end=True与save_strategy="epoch"结合时,若模型保存/加载过程中出现内存问题,也可能导致静默崩溃:

  • 暂时关闭load_best_model_at_end,测试早停是否正常
  • 确保output_dir有足够磁盘空间且权限正常

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 21:24:54