使用pyctcdecode进行Wav2Vec2波束搜索解码时遇max()参数错误求解
解决Wav2Vec2 + pyctcdecode波束搜索解码的错误问题
错误原因
pyctcdecode内部依赖numpy API处理数据,你传入的PyTorch张量的max()方法参数与numpy不兼容,触发类型错误。- 直接取
processor.tokenizer.get_vocab().keys()会得到乱序词汇表,与模型输出logits的维度顺序不匹配,后续解码逻辑会出现偏差。
修正后的代码
from pyctcdecode import build_ctcdecoder from transformers import Wav2Vec2ForCTC, Wav2Vec2Processor from torchaudio.utils import download_asset import torch import librosa # 加载模型与处理器 processor = Wav2Vec2Processor.from_pretrained("facebook/wav2vec2-base-960h") model = Wav2Vec2ForCTC.from_pretrained("facebook/wav2vec2-base-960h") # 下载测试音频 FILE_NAME = "tutorial-assets/Lab41-SRI-VOiCES-src-sp0307-ch127535-sg0042.wav" SPEECH_FILE = download_asset(FILE_NAME) # 加载并预处理音频 speech, sr = librosa.load(SPEECH_FILE, sr=16000) input_values = processor(speech, sampling_rate=16000, return_tensors="pt").input_values # 获取模型输出并转换为log概率 logits = model(input_values).logits # 按token id排序生成匹配模型输出的词汇表 vocab_dict = processor.tokenizer.get_vocab() sorted_vocab = sorted(vocab_dict.items(), key=lambda x: x[1]) vocabulary = [token for token, id in sorted_vocab] # 将PyTorch张量转换为numpy数组,适配pyctcdecode逻辑 log_probs = torch.nn.functional.log_softmax(logits[0], dim=-1).detach().numpy() # 构建解码器并执行波束搜索解码 decoder = build_ctcdecoder(vocabulary) # beam_width控制波束搜索宽度,值越大精度越高、速度越慢 text = decoder.decode(log_probs, beam_width=10) print(text)
关键修正点
- 词汇表排序:通过
sorted()按token id排序,确保词汇表顺序与模型输出logits的维度一一对应。 - 张量转numpy:使用
.detach().numpy()将PyTorch张量转换为numpy数组,适配pyctcdecode的内部处理逻辑。 - 明确softmax维度:在
log_softmax中指定dim=-1,避免维度歧义导致的计算错误。
内容的提问来源于stack exchange,提问作者user3668129
相关产品推荐
相关产品推荐

