mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-29 18:03:56 +08:00
fix:高并发下,共享vad里decoder变量扰动
This commit is contained in:
@@ -2,7 +2,8 @@
|
|||||||
"permissions": {
|
"permissions": {
|
||||||
"allow": [
|
"allow": [
|
||||||
"Bash(tree:*)",
|
"Bash(tree:*)",
|
||||||
"Bash(find:*)"
|
"Bash(find:*)",
|
||||||
|
"Bash(python3:*)"
|
||||||
],
|
],
|
||||||
"deny": [],
|
"deny": [],
|
||||||
"ask": []
|
"ask": []
|
||||||
|
|||||||
@@ -1119,6 +1119,14 @@ class ConnectionHandler:
|
|||||||
async def close(self, ws=None):
|
async def close(self, ws=None):
|
||||||
"""资源清理方法"""
|
"""资源清理方法"""
|
||||||
try:
|
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"):
|
if hasattr(self, "audio_buffer"):
|
||||||
self.audio_buffer.clear()
|
self.audio_buffer.clear()
|
||||||
|
|||||||
@@ -1,7 +1,8 @@
|
|||||||
import time
|
import time
|
||||||
|
import os
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
|
||||||
import opuslib_next
|
import opuslib_next
|
||||||
|
import onnxruntime
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from core.providers.vad.base import VADProviderBase
|
from core.providers.vad.base import VADProviderBase
|
||||||
|
|
||||||
@@ -12,16 +13,17 @@ logger = setup_logging()
|
|||||||
class VADProvider(VADProviderBase):
|
class VADProvider(VADProviderBase):
|
||||||
def __init__(self, config):
|
def __init__(self, config):
|
||||||
logger.bind(tag=TAG).info("SileroVAD", config)
|
logger.bind(tag=TAG).info("SileroVAD", config)
|
||||||
self.model, _ = torch.hub.load(
|
|
||||||
repo_or_dir=config["model_dir"],
|
model_path = os.path.join(
|
||||||
source="local",
|
config["model_dir"], "src", "silero_vad", "data", "silero_vad.onnx"
|
||||||
model="silero_vad",
|
)
|
||||||
force_reload=False,
|
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 = config.get("threshold", "0.5")
|
||||||
threshold_low = config.get("threshold_low", "0.2")
|
threshold_low = config.get("threshold_low", "0.2")
|
||||||
min_silence_duration_ms = config.get("min_silence_duration_ms", "1000")
|
min_silence_duration_ms = config.get("min_silence_duration_ms", "1000")
|
||||||
@@ -33,15 +35,25 @@ class VADProvider(VADProviderBase):
|
|||||||
int(min_silence_duration_ms) if min_silence_duration_ms else 1000
|
int(min_silence_duration_ms) if min_silence_duration_ms else 1000
|
||||||
)
|
)
|
||||||
|
|
||||||
# 至少要多少帧才算有语音
|
|
||||||
self.frame_window_threshold = 3
|
self.frame_window_threshold = 3
|
||||||
|
|
||||||
def __del__(self):
|
def _init_connection_state(self, conn):
|
||||||
if hasattr(self, 'decoder') and self.decoder is not None:
|
"""为连接初始化独立的 VAD 状态"""
|
||||||
try:
|
if not hasattr(conn, "_vad_opus_decoder"):
|
||||||
del self.decoder
|
conn._vad_opus_decoder = opuslib_next.Decoder(16000, 1)
|
||||||
except Exception:
|
if not hasattr(conn, "_vad_state"):
|
||||||
pass
|
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):
|
def is_vad(self, conn, opus_packet):
|
||||||
# 手动模式:直接返回True,不进行实时VAD检测,所有音频都缓存
|
# 手动模式:直接返回True,不进行实时VAD检测,所有音频都缓存
|
||||||
@@ -49,24 +61,32 @@ class VADProvider(VADProviderBase):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
try:
|
try:
|
||||||
pcm_frame = self.decoder.decode(opus_packet, 960)
|
self._init_connection_state(conn)
|
||||||
conn.client_audio_buffer.extend(pcm_frame) # 将新数据加入缓冲区
|
|
||||||
|
pcm_frame = conn._vad_opus_decoder.decode(opus_packet, 960)
|
||||||
|
conn.client_audio_buffer.extend(pcm_frame)
|
||||||
|
|
||||||
# 处理缓冲区中的完整帧(每次处理512采样点)
|
|
||||||
client_have_voice = False
|
client_have_voice = False
|
||||||
while len(conn.client_audio_buffer) >= 512 * 2:
|
while len(conn.client_audio_buffer) >= 512 * 2:
|
||||||
# 提取前512个采样点(1024字节)
|
|
||||||
chunk = conn.client_audio_buffer[: 512 * 2]
|
chunk = conn.client_audio_buffer[: 512 * 2]
|
||||||
conn.client_audio_buffer = 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_int16 = np.frombuffer(chunk, dtype=np.int16)
|
||||||
audio_float32 = audio_int16.astype(np.float32) / 32768.0
|
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)
|
||||||
|
|
||||||
# 检测语音活动
|
ort_inputs = {
|
||||||
with torch.no_grad():
|
"input": audio_input,
|
||||||
speech_prob = self.model(audio_tensor, 16000).item()
|
"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:
|
if speech_prob >= self.vad_threshold:
|
||||||
|
|||||||
Reference in New Issue
Block a user