From 8d2ba39ab861b0f0a921b324bcc466e37da59b92 Mon Sep 17 00:00:00 2001 From: hrz <1710360675@qq.com> Date: Tue, 27 May 2025 18:51:08 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BC=98=E5=8C=96=E5=8F=8C=E6=B5=81=E5=BC=8FTT?= =?UTF-8?q?S=E6=97=B6=E7=9A=84=E5=A3=B0=E9=9F=B3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main/xiaozhi-server/core/connection.py | 50 +++++++++++-------- .../xiaozhi-server/core/handle/helloHandle.py | 7 ++- .../xiaozhi-server/core/providers/tts/base.py | 10 +++- .../core/providers/tts/dto/dto.py | 7 +++ .../providers/tts/huoshan_double_stream.py | 27 +++++----- 5 files changed, 63 insertions(+), 38 deletions(-) diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py index ccf67fff..3243f18b 100644 --- a/main/xiaozhi-server/core/connection.py +++ b/main/xiaozhi-server/core/connection.py @@ -298,30 +298,36 @@ class ConnectionHandler: ) def _initialize_components(self): - """初始化组件""" - if self.config.get("prompt") is not None: - self.prompt = self.config["prompt"] - self.change_system_prompt(self.prompt) - self.logger.bind(tag=TAG).info( - f"初始化组件: prompt成功 {self.prompt[:50]}..." + try: + + """初始化组件""" + if self.config.get("prompt") is not None: + self.prompt = self.config["prompt"] + self.change_system_prompt(self.prompt) + self.logger.bind(tag=TAG).info( + f"初始化组件: prompt成功 {self.prompt[:50]}..." + ) + + """初始化本地组件""" + if self.vad is None: + self.vad = self._vad + if self.asr is None: + self.asr = self._asr + if self.tts is None: + self.tts = self._initialize_tts() + # 使用事件循环运行异步方法 + asyncio.run_coroutine_threadsafe( + self.tts.open_audio_channels(self), self.loop ) - """初始化本地组件""" - if self.vad is None: - self.vad = self._vad - if self.asr is None: - self.asr = self._asr - if self.tts is None: - self.tts = self._initialize_tts() - # 使用事件循环运行异步方法 - asyncio.run_coroutine_threadsafe(self.tts.open_audio_channels(self), self.loop) - - """加载记忆""" - self._initialize_memory() - """加载意图识别""" - self._initialize_intent() - """初始化上报线程""" - self._init_report_threads() + """加载记忆""" + self._initialize_memory() + """加载意图识别""" + self._initialize_intent() + """初始化上报线程""" + self._init_report_threads() + except Exception as e: + self.logger.bind(tag=TAG).error(f"实例化组件失败: {e}") def _init_report_threads(self): """初始化ASR和TTS上报线程""" diff --git a/main/xiaozhi-server/core/handle/helloHandle.py b/main/xiaozhi-server/core/handle/helloHandle.py index 422afe70..6fd40401 100644 --- a/main/xiaozhi-server/core/handle/helloHandle.py +++ b/main/xiaozhi-server/core/handle/helloHandle.py @@ -4,10 +4,9 @@ import json import random import shutil import asyncio -from core.providers.tts.dto.dto import ContentType -from core.providers.tts.dto.dto import SentenceType from core.handle.sendAudioHandle import send_stt_message from core.utils.util import remove_punctuation_and_length +from core.providers.tts.dto.dto import ContentType, InterfaceType TAG = __name__ @@ -40,6 +39,10 @@ async def checkWakeupWords(conn, text): enable_wakeup_words_response_cache = conn.config[ "enable_wakeup_words_response_cache" ] + """是否用的是非流式tts""" + if conn.tts and conn.tts.interface_type != InterfaceType.NON_STREAM: + return False + """是否开启唤醒词加速""" if not enable_wakeup_words_response_cache: return False diff --git a/main/xiaozhi-server/core/providers/tts/base.py b/main/xiaozhi-server/core/providers/tts/base.py index c44ba640..0b17eef5 100644 --- a/main/xiaozhi-server/core/providers/tts/base.py +++ b/main/xiaozhi-server/core/providers/tts/base.py @@ -13,7 +13,12 @@ from core.utils.tts import MarkdownCleaner from core.utils.output_counter import add_device_output from core.handle.reportHandle import enqueue_tts_report from core.handle.sendAudioHandle import sendAudioMessage -from core.providers.tts.dto.dto import TTSMessageDTO, SentenceType, ContentType +from core.providers.tts.dto.dto import ( + TTSMessageDTO, + SentenceType, + ContentType, + InterfaceType, +) import traceback @@ -24,10 +29,11 @@ logger = setup_logging() class TTSProviderBase(ABC): def __init__(self, config, delete_audio_file): + self.interface_type = InterfaceType.NON_STREAM self.conn = None self.tts_timeout = 10 self.delete_audio_file = delete_audio_file - self.output_file = config.get("output_dir") + self.output_file = config.get("output_dir", "tmp/") self.tts_text_queue = queue.Queue() self.tts_audio_queue = queue.Queue() self.tts_audio_first_sentence = True diff --git a/main/xiaozhi-server/core/providers/tts/dto/dto.py b/main/xiaozhi-server/core/providers/tts/dto/dto.py index a57a27de..b2317340 100644 --- a/main/xiaozhi-server/core/providers/tts/dto/dto.py +++ b/main/xiaozhi-server/core/providers/tts/dto/dto.py @@ -16,6 +16,13 @@ class ContentType(Enum): ACTION = "ACTION" # 动作内容 +class InterfaceType(Enum): + # 接口类型 + DUAL_STREAM = "DUAL_STREAM" # 双流式 + SINGLE_STREAM = "SINGLE_STREAM" # 单流式 + NON_STREAM = "NON_STREAM" # 非流式 + + class TTSMessageDTO: def __init__( self, diff --git a/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py b/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py index 1bd4f14d..702314fd 100644 --- a/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py +++ b/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py @@ -1,15 +1,16 @@ import os +import uuid +import json +import queue import asyncio import threading import traceback -import uuid -import json import websockets -from core.utils import opus_encoder_utils -import queue from config.logger import setup_logging +from core.utils import opus_encoder_utils +from core.utils.util import check_model_key from core.providers.tts.base import TTSProviderBase -from core.providers.tts.dto.dto import SentenceType, ContentType +from core.providers.tts.dto.dto import SentenceType, ContentType, InterfaceType TAG = __name__ logger = setup_logging() @@ -137,6 +138,7 @@ class Response: class TTSProvider(TTSProviderBase): def __init__(self, config, delete_audio_file): super().__init__(config, delete_audio_file) + self.interface_type = InterfaceType.DUAL_STREAM self.appId = config.get("appid") self.access_token = config.get("access_token") self.cluster = config.get("cluster") @@ -154,6 +156,7 @@ class TTSProvider(TTSProviderBase): self.opus_encoder = opus_encoder_utils.OpusEncoderUtils( sample_rate=16000, channels=1, frame_size_ms=60 ) + check_model_key("TTS", self.access_token) ################################################################################### # 火山双流式TTS重写父类的方法--开始 @@ -203,13 +206,6 @@ class TTSProvider(TTSProviderBase): ) if message.sentence_type == SentenceType.LAST: - for tts_file, text in self.before_stop_play_files: - if tts_file and os.path.exists(tts_file): - audio_datas = self._process_audio_file(tts_file) - self.tts_audio_queue.put( - (message.sentence_type, audio_datas, text) - ) - self.before_stop_play_files.clear() future = asyncio.run_coroutine_threadsafe( self.finish_session(self.conn.sentence_id), loop=self.conn.loop ) @@ -262,6 +258,13 @@ class TTSProvider(TTSProviderBase): logger.bind(tag=TAG).debug(f"句子结束~~{self.tts_text}") elif res.optional.event == EVENT_SessionFinished: logger.bind(tag=TAG).debug(f"会话结束~~") + for tts_file, text in self.before_stop_play_files: + if tts_file and os.path.exists(tts_file): + audio_datas = self._process_audio_file(tts_file) + self.tts_audio_queue.put( + (SentenceType.MIDDLE, audio_datas, text) + ) + self.before_stop_play_files.clear() self.tts_audio_queue.put((SentenceType.LAST, [], None)) continue except websockets.ConnectionClosed: