基于wav2vec2-base-960h训练的模型评估正常,加载推理乱码求助
基于facebook/wav2vec2-base-960h预训练的Wav2Vec2模型,训练阶段trainer.evaluate()输出结果正常(如下示例),但加载模型推理时输出乱码文本,可从以下几个方向排查:
reference: "شما امروز صبوری بفرمایین ثبت شده تا امروز با شما هماهنگی انجام بشه"
predicted: "شما امروز سبوری بفرمای سبز شده تا امروز با شما همهمنگی انجام باشه"
推理乱码示例:رچسصجپ هدثج یو تو یتنپ هر وغسهروغج سچ ثزتسه شتذس صمرجچو
1. 训练时Tokenizer参数错误
训练代码中Trainer的tokenizer参数传入了processor.feature_extractor,这是核心错误:
# 错误写法 trainer = Trainer( ... tokenizer=processor.feature_extractor, # 此处应为processor.tokenizer )
Wav2Vec2ForCTC需要绑定文本Tokenizer来处理标签词汇表,而非音频特征提取器。该错误会导致模型训练时未关联正确的波斯语词汇表,虽然评估阶段可能临时使用了正确的Tokenizer,但保存的模型未绑定对应词汇表,最终推理解码时出现乱码。
修正方案:
重新初始化Trainer,传入正确的Tokenizer:
trainer = Trainer( model=model, data_collator=data_collator, args=training_args, compute_metrics=compute_metrics, train_dataset=_common_voice_train, eval_dataset=_common_voice_test, tokenizer=processor.tokenizer, # 替换为processor.tokenizer )
重新训练后,务必将processor完整保存到模型目录:
processor.save_pretrained(save_dir)
2. 推理时Processor与模型词汇表不匹配
若训练后未保存processor,或推理时加载的processor是默认英文版本,会导致解码时用错误的词汇表映射ID,输出乱码。
验证与解决:
- 检查模型目录下是否包含
tokenizer_config.json、vocab.json等Tokenizer相关文件; - 加载后查看
processor.tokenizer.vocab,确认包含波斯语字符; - 确保推理时加载的
processor与训练时使用的是同一实例,而非重新从预训练模型加载。
3. 音频预处理不一致
训练与推理阶段的音频预处理流程必须完全一致,否则模型输入特征异常会导致解码乱码。
检查点:
- 确认训练时
data_collator是否对音频做了截断、补全、归一化等操作,推理时processor调用需保持相同参数(如padding=True、truncation=True); - 替换音频加载方式,避免librosa可能的格式问题:
import torchaudio audio_input, sample_rate = torchaudio.load("/content60_L4.wav") # 确保采样率转为16000 audio_input = torchaudio.functional.resample(audio_input, sample_rate, 16000) audio_input = audio_input.squeeze().numpy() # 转为单通道numpy数组
4. 模型加载的权重或精度问题
- 确认加载的是训练完成后的最终模型,而非中间检查点,可指定
local_files_only=True确保加载本地文件:
model = Wav2Vec2ForCTC.from_pretrained("/content/drive/MyDrive/model", local_files_only=True)
- 训练时开启了
fp16=True,推理时可尝试匹配精度:
# 若使用GPU,开启半精度 model = Wav2Vec2ForCTC.from_pretrained("/content/drive/MyDrive/model").to("cuda").half() # 或转为float32推理 model = Wav2Vec2ForCTC.from_pretrained("/content/drive/MyDrive/model").float()
内容的提问来源于stack exchange,提问作者miladjurablu

