优化双流式TTS时的声音

This commit is contained in:
hrz
2025-05-27 18:51:08 +08:00
parent 4260e5a1a7
commit 8d2ba39ab8
5 changed files with 63 additions and 38 deletions
@@ -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
@@ -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,
@@ -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: