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

使用HuggingFace训练WhisperSmall遇RuntimeError:重复反向传播问题求解

解决WhisperSmall训练时的RuntimeError问题

错误信息

RuntimeError: Trying to backward through the graph a second time (or directly access saved tensors after they have already been freed). Saved intermediate values of the graph are freed when you call .backward() or autograd.grad(). Specify retain_graph=True if you need to backward through the graph a second time or if you need to access saved tensors after calling backward.

用户代码片段

from transformers import WhisperForConditionalGeneration
model = WhisperForConditionalGeneration.from_pretrained("openai/whisper-small")

model.generation_config.language = "english"
model.generation_config.task = "transcribe"

model.config.use_cache= False
model.gradient_checkpointing_enable()
model_gradient_checkpointing= True

model.generation_config.forced_decoder_ids = None


from transformers import Seq2SeqTrainingArguments

training_args = Seq2SeqTrainingArguments(
    output_dir="./whisper-small-NER",  # change to a repo name of your choice
    per_device_train_batch_size=16,
    gradient_accumulation_steps=1,  # increase by 2x for every 2x decrease in batch size
    learning_rate=1e-5,
    warmup_steps=500,
    max_steps=5000,
    gradient_checkpointing=True,
    fp16=True,
    #fp16=False,
    eval_strategy="steps",
    per_device_eval_batch_size=8,
    predict_with_generate=True,
    #predict_with_generate=False,
    generation_max_length=225,
    save_steps=1000,
    eval_steps=1000,
    logging_steps=25,
    report_to=["tensorboard"],
    load_best_model_at_end=True,
    metric_for_best_model="wer",
    greater_is_better=False,
    push_to_hub=False,
)

from transformers import Seq2SeqTrainer

trainer = Seq2SeqTrainer(
    args=training_args,
    model=model,
    train_dataset=dataset_dict["train"],
    eval_dataset=dataset_dict["test"],
    data_collator=data_collator,
    compute_metrics=compute_metrics,
    tokenizer=processor.feature_extractor,
)

trainer.train()

解决方法

1. 修正梯度检查点与缓存配置

错误核心原因:开启梯度检查点(gradient_checkpointing=True)后,训练阶段会释放计算图的中间张量,但eval阶段开启predict_with_generate=True时,会再次尝试访问这些已释放的张量,导致冲突。

修改代码如下:

  • 移除手动调用的model.gradient_checkpointing_enable(),让Seq2SeqTrainingArguments中的gradient_checkpointing=True自动管理模型配置,Trainer会在训练/eval阶段自动切换合适的参数。
  • 删除无效的model_gradient_checkpointing= True变量(该行无实际作用)。
  • 不要手动设置model.config.use_cache=False,Trainer会在训练时自动关闭use_cache,eval时自动开启,避免配置冲突。

修改后的模型初始化代码:

from transformers import WhisperForConditionalGeneration
model = WhisperForConditionalGeneration.from_pretrained("openai/whisper-small")

model.generation_config.language = "english"
model.generation_config.task = "transcribe"
model.generation_config.forced_decoder_ids = None

2. 可选:拆分训练与eval生成

如果修改后仍有问题,可以暂时关闭predict_with_generate=False完成训练,之后再单独运行eval并开启generate计算指标:

training_args = Seq2SeqTrainingArguments(
    # 其他参数保持不变
    predict_with_generate=False,
)

训练完成后单独执行eval:

trainer.predict(test_dataset=dataset_dict["test"], predict_with_generate=True)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 19:45:16