From fd8296262e4213e9636eafd7cbeb8cad79c96278 Mon Sep 17 00:00:00 2001 From: hrz <1710360675@qq.com> Date: Tue, 24 Feb 2026 01:59:44 +0800 Subject: [PATCH] =?UTF-8?q?fix:=E9=AB=98=E5=B9=B6=E5=8F=91=E4=B8=8B?= =?UTF-8?q?=EF=BC=8C=E5=85=B1=E4=BA=ABvad=E9=87=8Cdecoder=E5=8F=98?= =?UTF-8?q?=E9=87=8F=E6=89=B0=E5=8A=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../.claude/settings.local.json | 3 +- main/xiaozhi-server/core/connection.py | 8 ++ .../core/providers/vad/silero.py | 74 ++++++++++++------- 3 files changed, 57 insertions(+), 28 deletions(-) diff --git a/main/xiaozhi-server/.claude/settings.local.json b/main/xiaozhi-server/.claude/settings.local.json index b9a1f5de..32b9ba94 100644 --- a/main/xiaozhi-server/.claude/settings.local.json +++ b/main/xiaozhi-server/.claude/settings.local.json @@ -2,7 +2,8 @@ "permissions": { "allow": [ "Bash(tree:*)", - "Bash(find:*)" + "Bash(find:*)", + "Bash(python3:*)" ], "deny": [], "ask": [] diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py index 6e8f5a8b..1d5d89c8 100644 --- a/main/xiaozhi-server/core/connection.py +++ b/main/xiaozhi-server/core/connection.py @@ -1119,6 +1119,14 @@ class ConnectionHandler: async def close(self, ws=None): """资源清理方法""" try: + # 清理 VAD 连接资源 + if ( + hasattr(self, "vad") + and self.vad + and hasattr(self.vad, "release_conn_resources") + ): + self.vad.release_conn_resources(self) + # 清理音频缓冲区 if hasattr(self, "audio_buffer"): self.audio_buffer.clear() diff --git a/main/xiaozhi-server/core/providers/vad/silero.py b/main/xiaozhi-server/core/providers/vad/silero.py index 81215681..8d3891dc 100644 --- a/main/xiaozhi-server/core/providers/vad/silero.py +++ b/main/xiaozhi-server/core/providers/vad/silero.py @@ -1,7 +1,8 @@ import time +import os import numpy as np -import torch import opuslib_next +import onnxruntime from config.logger import setup_logging from core.providers.vad.base import VADProviderBase @@ -12,16 +13,17 @@ logger = setup_logging() class VADProvider(VADProviderBase): def __init__(self, config): logger.bind(tag=TAG).info("SileroVAD", config) - self.model, _ = torch.hub.load( - repo_or_dir=config["model_dir"], - source="local", - model="silero_vad", - force_reload=False, + + model_path = os.path.join( + config["model_dir"], "src", "silero_vad", "data", "silero_vad.onnx" + ) + opts = onnxruntime.SessionOptions() + opts.inter_op_num_threads = 1 + opts.intra_op_num_threads = 1 + self.session = onnxruntime.InferenceSession( + model_path, providers=["CPUExecutionProvider"], sess_options=opts ) - self.decoder = opuslib_next.Decoder(16000, 1) - - # 处理空字符串的情况 threshold = config.get("threshold", "0.5") threshold_low = config.get("threshold_low", "0.2") min_silence_duration_ms = config.get("min_silence_duration_ms", "1000") @@ -33,40 +35,58 @@ class VADProvider(VADProviderBase): int(min_silence_duration_ms) if min_silence_duration_ms else 1000 ) - # 至少要多少帧才算有语音 self.frame_window_threshold = 3 - def __del__(self): - if hasattr(self, 'decoder') and self.decoder is not None: - try: - del self.decoder - except Exception: - pass + def _init_connection_state(self, conn): + """为连接初始化独立的 VAD 状态""" + if not hasattr(conn, "_vad_opus_decoder"): + conn._vad_opus_decoder = opuslib_next.Decoder(16000, 1) + if not hasattr(conn, "_vad_state"): + conn._vad_state = np.zeros((2, 1, 128), dtype=np.float32) + if not hasattr(conn, "_vad_context"): + conn._vad_context = np.zeros((1, 64), dtype=np.float32) + + def release_conn_resources(self, conn): + """释放连接的 VAD 资源(连接关闭时调用)""" + for attr in ("_vad_opus_decoder", "_vad_state", "_vad_context"): + if hasattr(conn, attr): + try: + delattr(conn, attr) + except Exception: + pass def is_vad(self, conn, opus_packet): # 手动模式:直接返回True,不进行实时VAD检测,所有音频都缓存 if conn.client_listen_mode == "manual": return True - - try: - pcm_frame = self.decoder.decode(opus_packet, 960) - conn.client_audio_buffer.extend(pcm_frame) # 将新数据加入缓冲区 - # 处理缓冲区中的完整帧(每次处理512采样点) + try: + self._init_connection_state(conn) + + pcm_frame = conn._vad_opus_decoder.decode(opus_packet, 960) + conn.client_audio_buffer.extend(pcm_frame) + client_have_voice = False while len(conn.client_audio_buffer) >= 512 * 2: - # 提取前512个采样点(1024字节) chunk = conn.client_audio_buffer[: 512 * 2] conn.client_audio_buffer = conn.client_audio_buffer[512 * 2 :] - # 转换为模型需要的张量格式 audio_int16 = np.frombuffer(chunk, dtype=np.int16) audio_float32 = audio_int16.astype(np.float32) / 32768.0 - audio_tensor = torch.from_numpy(audio_float32) + audio_input = np.concatenate( + [conn._vad_context, audio_float32.reshape(1, -1)], axis=1 + ).astype(np.float32) - # 检测语音活动 - with torch.no_grad(): - speech_prob = self.model(audio_tensor, 16000).item() + ort_inputs = { + "input": audio_input, + "state": conn._vad_state, + "sr": np.array(16000, dtype=np.int64), + } + out, state = self.session.run(None, ort_inputs) + + conn._vad_state = state + conn._vad_context = audio_input[:, -64:] + speech_prob = out.item() # 双阈值判断 if speech_prob >= self.vad_threshold: