使用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
相关产品推荐
相关产品推荐

