mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-26 17:13:54 +08:00
Merge pull request #1852 from xinnan-tech/py_fix_DoubaoStreamASR
修复关于声纹识别的若干问题
This commit is contained in:
@@ -24,7 +24,6 @@ from core.utils.modules_initialize import (
|
|||||||
initialize_asr,
|
initialize_asr,
|
||||||
)
|
)
|
||||||
from core.handle.reportHandle import report
|
from core.handle.reportHandle import report
|
||||||
from core.utils.modules_initialize import initialize_voiceprint
|
|
||||||
from core.providers.tts.default import DefaultTTS
|
from core.providers.tts.default import DefaultTTS
|
||||||
from concurrent.futures import ThreadPoolExecutor
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
from core.utils.dialogue import Message, Dialogue
|
from core.utils.dialogue import Message, Dialogue
|
||||||
@@ -39,6 +38,7 @@ from core.providers.tts.dto.dto import ContentType, TTSMessageDTO, SentenceType
|
|||||||
from config.logger import setup_logging, build_module_string, create_connection_logger
|
from config.logger import setup_logging, build_module_string, create_connection_logger
|
||||||
from config.manage_api_client import DeviceNotFoundException, DeviceBindException
|
from config.manage_api_client import DeviceNotFoundException, DeviceBindException
|
||||||
from core.utils.prompt_manager import PromptManager
|
from core.utils.prompt_manager import PromptManager
|
||||||
|
from core.utils.voiceprint_provider import VoiceprintProvider
|
||||||
|
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
@@ -109,6 +109,9 @@ class ConnectionHandler:
|
|||||||
self.memory = _memory
|
self.memory = _memory
|
||||||
self.intent = _intent
|
self.intent = _intent
|
||||||
|
|
||||||
|
# 为每个连接单独管理声纹识别
|
||||||
|
self.voiceprint_provider = None
|
||||||
|
|
||||||
# vad相关变量
|
# vad相关变量
|
||||||
self.client_audio_buffer = bytearray()
|
self.client_audio_buffer = bytearray()
|
||||||
self.client_have_voice = False
|
self.client_have_voice = False
|
||||||
@@ -348,6 +351,10 @@ class ConnectionHandler:
|
|||||||
self.vad = self._vad
|
self.vad = self._vad
|
||||||
if self.asr is None:
|
if self.asr is None:
|
||||||
self.asr = self._initialize_asr()
|
self.asr = self._initialize_asr()
|
||||||
|
|
||||||
|
# 初始化声纹识别
|
||||||
|
self._initialize_voiceprint()
|
||||||
|
|
||||||
# 打开语音识别通道
|
# 打开语音识别通道
|
||||||
asyncio.run_coroutine_threadsafe(
|
asyncio.run_coroutine_threadsafe(
|
||||||
self.asr.open_audio_channels(self), self.loop
|
self.asr.open_audio_channels(self), self.loop
|
||||||
@@ -416,17 +423,19 @@ class ConnectionHandler:
|
|||||||
# 因为远程ASR,涉及到websocket连接和接收线程,需要每个连接一个实例
|
# 因为远程ASR,涉及到websocket连接和接收线程,需要每个连接一个实例
|
||||||
asr = initialize_asr(self.config)
|
asr = initialize_asr(self.config)
|
||||||
|
|
||||||
# 动态初始化声纹识别功能
|
return asr
|
||||||
|
|
||||||
|
def _initialize_voiceprint(self):
|
||||||
|
"""为当前连接初始化声纹识别"""
|
||||||
try:
|
try:
|
||||||
success = initialize_voiceprint(asr, self.config)
|
voiceprint_config = self.config.get("voiceprint", {})
|
||||||
if success:
|
if voiceprint_config:
|
||||||
|
self.voiceprint_provider = VoiceprintProvider(voiceprint_config)
|
||||||
self.logger.bind(tag=TAG).info("声纹识别功能已在连接时动态启用")
|
self.logger.bind(tag=TAG).info("声纹识别功能已在连接时动态启用")
|
||||||
else:
|
else:
|
||||||
self.logger.bind(tag=TAG).info("声纹识别功能未启用或配置不完整")
|
self.logger.bind(tag=TAG).info("声纹识别功能未启用或配置不完整")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.logger.bind(tag=TAG).error(f"动态初始化声纹识别时发生错误: {str(e)}")
|
self.logger.bind(tag=TAG).warning(f"声纹识别初始化失败: {str(e)}")
|
||||||
|
|
||||||
return asr
|
|
||||||
|
|
||||||
def _initialize_private_config(self):
|
def _initialize_private_config(self):
|
||||||
"""如果是从配置文件获取,则进行二次实例化"""
|
"""如果是从配置文件获取,则进行二次实例化"""
|
||||||
|
|||||||
@@ -17,7 +17,6 @@ from core.handle.receiveAudioHandle import startToChat
|
|||||||
from core.handle.reportHandle import enqueue_asr_report
|
from core.handle.reportHandle import enqueue_asr_report
|
||||||
from core.utils.util import remove_punctuation_and_length
|
from core.utils.util import remove_punctuation_and_length
|
||||||
from core.handle.receiveAudioHandle import handleAudioMessage
|
from core.handle.receiveAudioHandle import handleAudioMessage
|
||||||
from core.utils.voiceprint_provider import VoiceprintProvider
|
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -25,13 +24,7 @@ logger = setup_logging()
|
|||||||
|
|
||||||
class ASRProviderBase(ABC):
|
class ASRProviderBase(ABC):
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.voiceprint_provider = None
|
pass
|
||||||
|
|
||||||
def init_voiceprint(self, voiceprint_config: dict):
|
|
||||||
"""初始化声纹识别"""
|
|
||||||
if voiceprint_config:
|
|
||||||
self.voiceprint_provider = VoiceprintProvider(voiceprint_config)
|
|
||||||
logger.bind(tag=TAG).info("声纹识别模块已初始化")
|
|
||||||
|
|
||||||
# 打开音频通道
|
# 打开音频通道
|
||||||
async def open_audio_channels(self, conn):
|
async def open_audio_channels(self, conn):
|
||||||
@@ -94,7 +87,8 @@ class ASRProviderBase(ABC):
|
|||||||
|
|
||||||
# 预先准备WAV数据
|
# 预先准备WAV数据
|
||||||
wav_data = None
|
wav_data = None
|
||||||
if self.voiceprint_provider and combined_pcm_data:
|
# 使用连接的声纹识别提供者
|
||||||
|
if conn.voiceprint_provider and combined_pcm_data:
|
||||||
wav_data = self._pcm_to_wav(combined_pcm_data)
|
wav_data = self._pcm_to_wav(combined_pcm_data)
|
||||||
|
|
||||||
|
|
||||||
@@ -102,7 +96,6 @@ class ASRProviderBase(ABC):
|
|||||||
def run_asr():
|
def run_asr():
|
||||||
start_time = time.monotonic()
|
start_time = time.monotonic()
|
||||||
try:
|
try:
|
||||||
import asyncio
|
|
||||||
loop = asyncio.new_event_loop()
|
loop = asyncio.new_event_loop()
|
||||||
asyncio.set_event_loop(loop)
|
asyncio.set_event_loop(loop)
|
||||||
try:
|
try:
|
||||||
@@ -123,14 +116,13 @@ class ASRProviderBase(ABC):
|
|||||||
def run_voiceprint():
|
def run_voiceprint():
|
||||||
if not wav_data:
|
if not wav_data:
|
||||||
return None
|
return None
|
||||||
start_time = time.monotonic()
|
|
||||||
try:
|
try:
|
||||||
import asyncio
|
|
||||||
loop = asyncio.new_event_loop()
|
loop = asyncio.new_event_loop()
|
||||||
asyncio.set_event_loop(loop)
|
asyncio.set_event_loop(loop)
|
||||||
try:
|
try:
|
||||||
|
# 使用连接的声纹识别提供者
|
||||||
result = loop.run_until_complete(
|
result = loop.run_until_complete(
|
||||||
self.voiceprint_provider.identify_speaker(wav_data, conn.session_id)
|
conn.voiceprint_provider.identify_speaker(wav_data, conn.session_id)
|
||||||
)
|
)
|
||||||
return result
|
return result
|
||||||
finally:
|
finally:
|
||||||
@@ -145,7 +137,7 @@ class ASRProviderBase(ABC):
|
|||||||
with concurrent.futures.ThreadPoolExecutor(max_workers=2) as thread_executor:
|
with concurrent.futures.ThreadPoolExecutor(max_workers=2) as thread_executor:
|
||||||
asr_future = thread_executor.submit(run_asr)
|
asr_future = thread_executor.submit(run_asr)
|
||||||
|
|
||||||
if self.voiceprint_provider and wav_data:
|
if conn.voiceprint_provider and wav_data:
|
||||||
voiceprint_future = thread_executor.submit(run_voiceprint)
|
voiceprint_future = thread_executor.submit(run_voiceprint)
|
||||||
|
|
||||||
# 等待两个线程都完成
|
# 等待两个线程都完成
|
||||||
@@ -157,7 +149,6 @@ class ASRProviderBase(ABC):
|
|||||||
asr_result = asr_future.result(timeout=15)
|
asr_result = asr_future.result(timeout=15)
|
||||||
results = {"asr": asr_result, "voiceprint": None}
|
results = {"asr": asr_result, "voiceprint": None}
|
||||||
|
|
||||||
parallel_execution_time = time.monotonic() - parallel_start_time
|
|
||||||
|
|
||||||
# 处理结果
|
# 处理结果
|
||||||
raw_text, file_path = results.get("asr", ("", None))
|
raw_text, file_path = results.get("asr", ("", None))
|
||||||
|
|||||||
@@ -57,6 +57,16 @@ class ASRProvider(ASRProviderBase):
|
|||||||
conn.asr_audio.append(audio)
|
conn.asr_audio.append(audio)
|
||||||
conn.asr_audio = conn.asr_audio[-10:]
|
conn.asr_audio = conn.asr_audio[-10:]
|
||||||
|
|
||||||
|
# 存储音频数据
|
||||||
|
if not hasattr(conn, 'asr_audio_for_voiceprint'):
|
||||||
|
conn.asr_audio_for_voiceprint = []
|
||||||
|
conn.asr_audio_for_voiceprint.append(audio)
|
||||||
|
|
||||||
|
# 当没有音频数据时处理完整语音片段
|
||||||
|
if not audio and len(conn.asr_audio_for_voiceprint) > 0:
|
||||||
|
await self.handle_voice_stop(conn, conn.asr_audio_for_voiceprint)
|
||||||
|
conn.asr_audio_for_voiceprint = []
|
||||||
|
|
||||||
# 如果本次有声音,且之前没有建立连接
|
# 如果本次有声音,且之前没有建立连接
|
||||||
if audio_have_voice and self.asr_ws is None and not self.is_processing:
|
if audio_have_voice and self.asr_ws is None and not self.is_processing:
|
||||||
try:
|
try:
|
||||||
@@ -148,6 +158,8 @@ class ASRProvider(ASRProviderBase):
|
|||||||
async def _forward_asr_results(self, conn):
|
async def _forward_asr_results(self, conn):
|
||||||
try:
|
try:
|
||||||
while self.asr_ws and not conn.stop_event.is_set():
|
while self.asr_ws and not conn.stop_event.is_set():
|
||||||
|
# 获取当前连接的音频数据
|
||||||
|
audio_data = getattr(conn, 'asr_audio_for_voiceprint', [])
|
||||||
try:
|
try:
|
||||||
response = await self.asr_ws.recv()
|
response = await self.asr_ws.recv()
|
||||||
result = self.parse_response(response)
|
result = self.parse_response(response)
|
||||||
@@ -171,7 +183,8 @@ class ASRProvider(ASRProviderBase):
|
|||||||
logger.bind(tag=TAG).error(f"识别文本:空")
|
logger.bind(tag=TAG).error(f"识别文本:空")
|
||||||
self.text = ""
|
self.text = ""
|
||||||
conn.reset_vad_states()
|
conn.reset_vad_states()
|
||||||
await self.handle_voice_stop(conn, None)
|
if len(audio_data) > 15: # 确保有足够音频数据
|
||||||
|
await self.handle_voice_stop(conn, audio_data)
|
||||||
break
|
break
|
||||||
|
|
||||||
for utterance in utterances:
|
for utterance in utterances:
|
||||||
@@ -181,7 +194,8 @@ class ASRProvider(ASRProviderBase):
|
|||||||
f"识别到文本: {self.text}"
|
f"识别到文本: {self.text}"
|
||||||
)
|
)
|
||||||
conn.reset_vad_states()
|
conn.reset_vad_states()
|
||||||
await self.handle_voice_stop(conn, None)
|
if len(audio_data) > 15: # 确保有足够音频数据
|
||||||
|
await self.handle_voice_stop(conn, audio_data)
|
||||||
break
|
break
|
||||||
elif "error" in payload:
|
elif "error" in payload:
|
||||||
error_msg = payload.get("error", "未知错误")
|
error_msg = payload.get("error", "未知错误")
|
||||||
@@ -208,6 +222,13 @@ class ASRProvider(ASRProviderBase):
|
|||||||
await self.asr_ws.close()
|
await self.asr_ws.close()
|
||||||
self.asr_ws = None
|
self.asr_ws = None
|
||||||
self.is_processing = False
|
self.is_processing = False
|
||||||
|
if conn:
|
||||||
|
if hasattr(conn, 'asr_audio_for_voiceprint'):
|
||||||
|
conn.asr_audio_for_voiceprint = []
|
||||||
|
if hasattr(conn, 'asr_audio'):
|
||||||
|
conn.asr_audio = []
|
||||||
|
if hasattr(conn, 'has_valid_voice'):
|
||||||
|
conn.has_valid_voice = False
|
||||||
|
|
||||||
def stop_ws_connection(self):
|
def stop_ws_connection(self):
|
||||||
if self.asr_ws:
|
if self.asr_ws:
|
||||||
@@ -349,3 +370,12 @@ class ASRProvider(ASRProviderBase):
|
|||||||
pass
|
pass
|
||||||
self.forward_task = None
|
self.forward_task = None
|
||||||
self.is_processing = False
|
self.is_processing = False
|
||||||
|
# 清理所有连接的音频缓冲区
|
||||||
|
if hasattr(self, '_connections'):
|
||||||
|
for conn in self._connections.values():
|
||||||
|
if hasattr(conn, 'asr_audio_for_voiceprint'):
|
||||||
|
conn.asr_audio_for_voiceprint = []
|
||||||
|
if hasattr(conn, 'asr_audio'):
|
||||||
|
conn.asr_audio = []
|
||||||
|
if hasattr(conn, 'has_valid_voice'):
|
||||||
|
conn.has_valid_voice = False
|
||||||
|
|||||||
Reference in New Issue
Block a user