基于CTC解码器的实时ASR逐块无损失高效实现问询
实时CTC-ASR无损失逐块处理方案
要解决实时逐块处理音频时的信号丢失问题,核心是利用上下文缓存机制——Wav2Vec2这类声学模型依赖卷积的上下文感受野,直接切分块处理会丢失块边缘的上下文信息,导致识别误差或输出不连续。下面是具体实现方案:
核心思路
- 计算模型的输入输出下采样率:Wav2Vec2_ASR_BASE_10M的总下采样率为320(即每320个输入采样点对应1个输出帧),这决定了每段输入音频对应的输出帧数量,以及需要保留的上下文长度。
- 维护上下文缓存:每次处理新音频块前,将前一块的尾部上下文(长度等于模型感受野对应的输入采样数)与当前块拼接,确保模型能获取足够的上下文信息。
- 提取有效输出:模型处理拼接后的音频后,只保留当前块对应的输出帧,避免重复处理上下文部分的结果。
修正后的代码实现
import torch import torchaudio import pyaudio # 加载模型 bundle = torchaudio.pipelines.WAV2VEC2_ASR_BASE_10M acoustic_model = bundle.get_model().to("cuda" if torch.cuda.is_available() else "cpu") labels = bundle.get_labels() sample_rate = bundle.sample_rate # 初始化PyAudio流 mic = pyaudio.PyAudio() block_size = 4096 # 可根据延迟需求调整,比如2048/1024 stream = mic.open( format=pyaudio.paFloat32, channels=1, rate=sample_rate, input=True, frames_per_buffer=block_size ) # 模型参数计算:下采样率、需要保留的上下文长度 test_input = torch.randn(1, sample_rate).to(acoustic_model.device) test_emission, _ = acoustic_model(test_input) downsample_rate = sample_rate // test_emission.shape[1] # 上下文长度设为2倍下采样率,确保覆盖模型卷积感受野 context_len = 2 * downsample_rate # 初始化上下文缓存和总emission context_cache = torch.empty(0, device=acoustic_model.device) total_emission = torch.empty(0, test_emission.shape[1], test_emission.shape[2], device=acoustic_model.device) # 实时处理循环 try: while True: # 读取音频块 data_bytes = stream.read(block_size) current_block = torch.frombuffer( data_bytes, dtype=torch.float32, requires_grad=False ).to(acoustic_model.device) # 拼接上下文缓存与当前块 input_audio = torch.cat([context_cache, current_block]) # 添加batch维度 input_audio = input_audio.unsqueeze(0) # 模型推理 emission, _ = acoustic_model(input_audio) # 提取当前块对应的输出帧:跳过上下文缓存对应的帧 context_frames = len(context_cache) // downsample_rate current_emission = emission[:, context_frames:, :] # 更新总emission用于CTC解码 total_emission = torch.cat([total_emission, current_emission], dim=1) # 更新上下文缓存:保留当前块的尾部context_len个采样点 context_cache = current_block[-context_len:] if len(current_block) >= context_len else current_block # 实时解码 decoder = ctc_decoder( lexicon=None, tokens=labels, beam_size=10, ) result = decoder(total_emission) transcript = " ".join(result[0][0].words).strip() print(f"实时转录: {transcript}") except KeyboardInterrupt: print("停止录音") finally: stream.stop_stream() stream.close() mic.terminate()
关键细节说明
- 修正输入形状:原代码中
torch.stack((data_t, data_t))是错误操作,模型仅需单batch单通道输入,用unsqueeze(0)添加batch维度即可,避免不必要的计算浪费。 - GPU加速:将模型和张量移到GPU可大幅降低推理耗时,这是实时处理的核心优化点。
- 块大小调整:
block_size越小延迟越低,但推理频率越高;可根据硬件性能调整为2048或1024。 - 增量解码平衡:若追求极致低延迟,可仅对
current_emission解码,但CTC的全局特性可能导致结果不稳定;更稳妥的方式是维护total_emission,每次解码全部历史帧,牺牲少量延迟换取准确率。
更简洁的替代方案:流式API
TorchAudio的Wav2Vec2ForCTC支持原生流式推理,自动处理上下文缓存,无需手动维护:
acoustic_model = bundle.get_model(streaming=True).to(device) # 初始化流式状态 state = acoustic_model.get_streaming_state() # 每次处理块时传入state emission, state = acoustic_model(input_audio, state=state)
内容的提问来源于stack exchange,提问作者rumnen
相关产品推荐
相关产品推荐

