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

使用pyctcdecode进行Wav2Vec2波束搜索解码时遇max()参数错误求解

解决Wav2Vec2 + pyctcdecode波束搜索解码的错误问题

错误原因

  1. pyctcdecode内部依赖numpy API处理数据,你传入的PyTorch张量的max()方法参数与numpy不兼容,触发类型错误。
  2. 直接取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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 04:30:26