fix: 播放咋音问题

This commit is contained in:
lizhongxiang
2025-04-09 14:32:32 +08:00
parent 4c81585e32
commit 8bb888c58e
5 changed files with 169 additions and 67 deletions
@@ -28,7 +28,7 @@ async def sendAudioMessage(conn, ttsMessageDTO: TTSMessageDTO):
# 流控参数优化 # 流控参数优化
original_frame_duration = 60 # 原始帧时长(毫秒) original_frame_duration = 60 # 原始帧时长(毫秒)
adjusted_frame_duration = int(original_frame_duration * 0.8) # 缩短20% adjusted_frame_duration = int(original_frame_duration * 1) # 缩短20%
total_frames = len(ttsMessageDTO.content) # 获取总帧数 total_frames = len(ttsMessageDTO.content) # 获取总帧数
compensation = ( compensation = (
total_frames * (original_frame_duration - adjusted_frame_duration) / 1000 total_frames * (original_frame_duration - adjusted_frame_duration) / 1000
@@ -54,8 +54,8 @@ async def sendAudioMessage(conn, ttsMessageDTO: TTSMessageDTO):
play_position += adjusted_frame_duration # 使用调整后的帧时长 play_position += adjusted_frame_duration # 使用调整后的帧时长
# 补偿因加速损失的时长 # 补偿因加速损失的时长
if compensation > 0: # if compensation > 0:
await asyncio.sleep(compensation) # await asyncio.sleep(compensation)
if SentenceType.SENTENCE_END == ttsMessageDTO.sentence_type: if SentenceType.SENTENCE_END == ttsMessageDTO.sentence_type:
logger.bind(tag=TAG).info(f"发送最后一段语音: {ttsMessageDTO.tts_finish_text}") logger.bind(tag=TAG).info(f"发送最后一段语音: {ttsMessageDTO.tts_finish_text}")
await send_tts_message(conn, "sentence_end", ttsMessageDTO.tts_finish_text) await send_tts_message(conn, "sentence_end", ttsMessageDTO.tts_finish_text)
@@ -0,0 +1,137 @@
"""
Opus编码工具类
将PCM音频数据编码为Opus格式
"""
import array
import logging
import traceback
import numpy as np
from typing import List, Optional
from opuslib import Encoder
from opuslib import constants
class OpusEncoderUtils:
"""PCM到Opus的编码器"""
def __init__(self, sample_rate: int, channels: int, frame_size_ms: int):
"""
初始化Opus编码器
Args:
sample_rate: 采样率 (Hz)
channels: 通道数 (1=单声道, 2=立体声)
frame_size_ms: 帧大小 (毫秒)
"""
self.sample_rate = sample_rate
self.channels = channels
self.frame_size_ms = frame_size_ms
# 计算每帧样本数 = 采样率 * 帧大小(毫秒) / 1000
self.frame_size = (sample_rate * frame_size_ms) // 1000
# 总帧大小 = 每帧样本数 * 通道数
self.total_frame_size = self.frame_size * channels
# 比特率和复杂度设置
self.bitrate = 24000 # bps
self.complexity = 10 # 最高质量
# 缓冲区初始化为空
self.buffer = np.array([], dtype=np.int16)
try:
# 创建Opus编码器
self.encoder = Encoder(
sample_rate,
channels,
constants.APPLICATION_AUDIO # 音频优化模式
)
self.encoder.bitrate = self.bitrate
self.encoder.complexity = self.complexity
self.encoder.signal = constants.SIGNAL_VOICE # 语音信号优化
except Exception as e:
logging.error(f"初始化Opus编码器失败: {e}")
raise RuntimeError("初始化失败") from e
def reset_state(self):
"""重置编码器状态"""
self.encoder.reset_state()
self.buffer = np.array([], dtype=np.int16)
def encode_pcm_to_opus(self, pcm_data: bytes, end_of_stream: bool) -> List[bytes]:
"""
将PCM数据编码为Opus格式
Args:
pcm_data: PCM字节数据
end_of_stream: 是否为流的结束
Returns:
Opus数据包列表
"""
# 将字节数据转换为short数组
new_samples = self._convert_bytes_to_shorts(pcm_data)
# 校验PCM数据
self._validate_pcm_data(new_samples)
# 将新数据追加到缓冲区
self.buffer = np.append(self.buffer, new_samples)
opus_packets = []
offset = 0
# 处理所有完整帧
while offset <= len(self.buffer) - self.total_frame_size:
frame = self.buffer[offset:offset + self.total_frame_size]
output = self._encode(frame)
if output:
opus_packets.append(output)
offset += self.total_frame_size
# 保留未处理的样本
self.buffer = self.buffer[offset:]
# 流结束时处理剩余数据
if end_of_stream and len(self.buffer) > 0:
# 创建最后一帧并用0填充
last_frame = np.zeros(self.total_frame_size, dtype=np.int16)
last_frame[:len(self.buffer)] = self.buffer
output = self._encode(last_frame)
if output:
opus_packets.append(output)
self.buffer = np.array([], dtype=np.int16)
return opus_packets
def _encode(self, frame: np.ndarray) -> Optional[bytes]:
"""编码一帧音频数据"""
try:
# 将numpy数组转换为bytes
frame_bytes = frame.tobytes()
# opuslib要求输入字节数必须是channels*2的倍数
encoded = self.encoder.encode(frame_bytes, self.frame_size)
return encoded
except Exception as e:
logging.error(f"Opus编码失败: {e}")
traceback.print_exc()
return None
def _convert_bytes_to_shorts(self, bytes_data: bytes) -> np.ndarray:
"""将字节数组转换为short数组 (16位PCM)"""
# 假设输入是小端字节序的16位PCM
return np.frombuffer(bytes_data, dtype=np.int16)
def _validate_pcm_data(self, pcm_shorts: np.ndarray) -> None:
"""验证PCM数据是否有效"""
# 16位PCM数据范围是 -32768 到 32767
if np.any((pcm_shorts < -32768) | (pcm_shorts > 32767)):
invalid_samples = pcm_shorts[(pcm_shorts < -32768) | (pcm_shorts > 32767)]
logging.warning(f"发现无效PCM样本: {invalid_samples[:5]}...")
# 在实际应用中可以选择裁剪而不是抛出异常
# np.clip(pcm_shorts, -32768, 32767, out=pcm_shorts)
def close(self):
"""关闭编码器并释放资源"""
# opuslib没有明确的关闭方法,Python的垃圾回收会处理
pass
+8 -45
View File
@@ -16,6 +16,7 @@ import opuslib_next
from pydub import AudioSegment from pydub import AudioSegment
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from core.utils import textUtils from core.utils import textUtils
from core.opus import opus_encoder_utils
import queue import queue
from core.providers.tts.dto.dto import MsgType, TTSMessageDTO, SentenceType from core.providers.tts.dto.dto import MsgType, TTSMessageDTO, SentenceType
@@ -33,6 +34,7 @@ class TTSProviderBase(ABC):
self.tts_audio_queue = queue.Queue() self.tts_audio_queue = queue.Queue()
self.enable_two_way = False self.enable_two_way = False
self.stop_event = threading.Event() self.stop_event = threading.Event()
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(sample_rate=16000, channels=1, frame_size_ms=60)
self.tts_text_buff = [] self.tts_text_buff = []
self.punctuations = ( self.punctuations = (
@@ -70,25 +72,18 @@ class TTSProviderBase(ABC):
) )
tts_priority.start() tts_priority.start()
async def stop_listen_resource(self):
"""资源清理方法"""
self.stop_event.set()
self.tts_text_queue = None
self.tts_audio_queue = None
gc.collect() # 强制执行垃圾回收
async def close(self): async def close(self):
pass self.stop_event
def _get_segment_text(self): def _get_segment_text(self):
# 合并当前全部文本并处理未分割部分 # 合并当前全部文本并处理未分割部分
full_text = "".join(self.tts_text_buff) full_text = "".join(self.tts_text_buff)
current_text = full_text[self.processed_chars :] # 从未处理的位置开始 current_text = full_text[self.processed_chars:] # 从未处理的位置开始
last_punct_pos = -1 last_punct_pos = -1
for punct in self.punctuations: for punct in self.punctuations:
pos = current_text.rfind(punct) pos = current_text.rfind(punct)
if (pos != -1 and last_punct_pos == -1) or ( if (pos != -1 and last_punct_pos == -1) or (
pos != -1 and pos < last_punct_pos pos != -1 and pos < last_punct_pos
): ):
last_punct_pos = pos last_punct_pos = pos
if last_punct_pos != -1: if last_punct_pos != -1:
@@ -118,7 +113,7 @@ class TTSProviderBase(ABC):
async def finish_session(self, session_id): async def finish_session(self, session_id):
pass pass
def tts_one_sentence(self,conn, text, u_id=None): def tts_one_sentence(self, conn, text, u_id=None):
if not u_id: if not u_id:
u_id = str(uuid.uuid4()).replace("-", "") u_id = str(uuid.uuid4()).replace("-", "")
conn.u_id = u_id conn.u_id = u_id
@@ -214,7 +209,6 @@ class TTSProviderBase(ABC):
) )
self.active_tasks.add(future) self.active_tasks.add(future)
if self.active_tasks: if self.active_tasks:
async def wrap_future(future): async def wrap_future(future):
return await asyncio.wrap_future(future) return await asyncio.wrap_future(future)
@@ -315,7 +309,7 @@ class TTSProviderBase(ABC):
# 按帧处理所有音频数据(包括最后一帧可能补零) # 按帧处理所有音频数据(包括最后一帧可能补零)
for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample
# 获取当前帧的二进制数据 # 获取当前帧的二进制数据
chunk = raw_data[i : i + frame_size * 2] chunk = raw_data[i: i + frame_size * 2]
# 如果最后一帧不足,补零 # 如果最后一帧不足,补零
if len(chunk) < frame_size * 2: if len(chunk) < frame_size * 2:
@@ -341,36 +335,5 @@ class TTSProviderBase(ABC):
return audio return audio
def wav_to_opus_data_audio_raw(self, raw_data_var, is_end=False): def wav_to_opus_data_audio_raw(self, raw_data_var, is_end=False):
raw_data = self.last_to_opus_raw + raw_data_var opus_datas = self.opus_encoder.encode_pcm_to_opus(raw_data_var, is_end)
self.last_to_opus_raw = b""
# 初始化Opus编码器
encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO)
# 编码参数
frame_duration = 60 # 60ms per frame
frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame
opus_datas = []
# 按帧处理所有音频数据(包括最后一帧可能补零)
for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample
# 获取当前帧的二进制数据
chunk = raw_data[i : i + frame_size * 2]
# 如果最后一帧不足,补零
# 缓存记录一下
if len(chunk) < frame_size * 2 and not is_end:
logger.bind(tag=TAG).info("如果最后一帧不足,缓存记录一下")
self.last_to_opus_raw = chunk
break
if len(chunk) < frame_size * 2 and is_end:
logger.bind(tag=TAG).info("是最后一句了,补零")
chunk += b"\x00" * (frame_size * 2 - len(chunk))
# 转换为numpy数组处理
np_frame = np.frombuffer(chunk, dtype=np.int16)
# 编码Opus数据
opus_data = encoder.encode(np_frame.tobytes(), frame_size)
opus_datas.append(opus_data)
return opus_datas return opus_datas
@@ -356,6 +356,7 @@ class TTSProvider(TTSProviderBase):
await super().reset() await super().reset()
async def close(self): async def close(self):
super().close()
"""资源清理方法""" """资源清理方法"""
await self.finish_connection() await self.finish_connection()
await self.ws.close() await self.ws.close()
@@ -373,7 +374,7 @@ class TTSProvider(TTSProviderBase):
async def _start_monitor_tts_response(self): async def _start_monitor_tts_response(self):
chunk_total = b"" chunk_total = b""
while True: while not self.stop_event.is_set():
try: try:
msg = await self.ws.recv() # 确保 `recv()` 运行在同一个 event loop msg = await self.ws.recv() # 确保 `recv()` 运行在同一个 event loop
res = self.parser_response(msg) res = self.parser_response(msg)
@@ -385,6 +386,7 @@ class TTSProvider(TTSProviderBase):
): ):
logger.bind(tag=TAG).info(f"推送数据到队列里面~~") logger.bind(tag=TAG).info(f"推送数据到队列里面~~")
opus_datas = self.wav_to_opus_data_audio_raw(res.payload) opus_datas = self.wav_to_opus_data_audio_raw(res.payload)
logger.bind(tag=TAG).info(f"推送数据到队列里面帧数~~{len(opus_datas)}")
self.tts_audio_queue.put( self.tts_audio_queue.put(
TTSMessageDTO( TTSMessageDTO(
u_id=self.u_id, u_id=self.u_id,
+18 -18
View File
@@ -12,7 +12,7 @@ class WebSocketServer:
def __init__(self, config: dict): def __init__(self, config: dict):
self.config = config self.config = config
self.logger = setup_logging() self.logger = setup_logging()
self._vad, self._asr, self._llm, self._tts, self._memory, self.intent = ( self._vad, self._asr, self._llm, self._memory, self.intent = (
self._create_processing_instances() self._create_processing_instances()
) )
self.active_connections = set() # 添加全局连接记录 self.active_connections = set() # 添加全局连接记录
@@ -22,7 +22,7 @@ class WebSocketServer:
"Memory", "nomem" "Memory", "nomem"
) # 默认使用nomem ) # 默认使用nomem
has_memory_cfg = ( has_memory_cfg = (
self.config.get("Memory") and memory_cls_name in self.config["Memory"] self.config.get("Memory") and memory_cls_name in self.config["Memory"]
) )
memory_cfg = self.config["Memory"][memory_cls_name] if has_memory_cfg else {} memory_cfg = self.config["Memory"][memory_cls_name] if has_memory_cfg else {}
@@ -36,7 +36,7 @@ class WebSocketServer:
( (
self.config["selected_module"]["ASR"] self.config["selected_module"]["ASR"]
if not "type" if not "type"
in self.config["ASR"][self.config["selected_module"]["ASR"]] in self.config["ASR"][self.config["selected_module"]["ASR"]]
else self.config["ASR"][self.config["selected_module"]["ASR"]][ else self.config["ASR"][self.config["selected_module"]["ASR"]][
"type" "type"
] ]
@@ -48,31 +48,19 @@ class WebSocketServer:
( (
self.config["selected_module"]["LLM"] self.config["selected_module"]["LLM"]
if not "type" if not "type"
in self.config["LLM"][self.config["selected_module"]["LLM"]] in self.config["LLM"][self.config["selected_module"]["LLM"]]
else self.config["LLM"][self.config["selected_module"]["LLM"]][ else self.config["LLM"][self.config["selected_module"]["LLM"]][
"type" "type"
] ]
), ),
self.config["LLM"][self.config["selected_module"]["LLM"]], self.config["LLM"][self.config["selected_module"]["LLM"]],
), ),
tts.create_instance(
(
self.config["selected_module"]["TTS"]
if not "type"
in self.config["TTS"][self.config["selected_module"]["TTS"]]
else self.config["TTS"][self.config["selected_module"]["TTS"]][
"type"
]
),
self.config["TTS"][self.config["selected_module"]["TTS"]],
self.config["delete_audio"],
),
memory.create_instance(memory_cls_name, memory_cfg), memory.create_instance(memory_cls_name, memory_cfg),
intent.create_instance( intent.create_instance(
( (
self.config["selected_module"]["Intent"] self.config["selected_module"]["Intent"]
if not "type" if not "type"
in self.config["Intent"][self.config["selected_module"]["Intent"]] in self.config["Intent"][self.config["selected_module"]["Intent"]]
else self.config["Intent"][ else self.config["Intent"][
self.config["selected_module"]["Intent"] self.config["selected_module"]["Intent"]
]["type"] ]["type"]
@@ -104,12 +92,24 @@ class WebSocketServer:
async def _handle_connection(self, websocket): async def _handle_connection(self, websocket):
"""处理新连接,每次创建独立的ConnectionHandler""" """处理新连接,每次创建独立的ConnectionHandler"""
# 创建ConnectionHandler时传入当前server实例 # 创建ConnectionHandler时传入当前server实例
_tts = tts.create_instance(
(
self.config["selected_module"]["TTS"]
if not "type"
in self.config["TTS"][self.config["selected_module"]["TTS"]]
else self.config["TTS"][self.config["selected_module"]["TTS"]][
"type"
]
),
self.config["TTS"][self.config["selected_module"]["TTS"]],
self.config["delete_audio"],
)
handler = ConnectionHandler( handler = ConnectionHandler(
self.config, self.config,
self._vad, self._vad,
self._asr, self._asr,
self._llm, self._llm,
self._tts, _tts,
self._memory, self._memory,
self.intent, self.intent,
) )