Files
aistream-test/python/tts_client-本地tts.py
T
2025-12-01 21:00:54 +08:00

551 lines
20 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import asyncio
import json
import websockets
import numpy as np
import sounddevice as sd
from typing import Optional, Callable, Dict, Any, Coroutine, List
from dataclasses import dataclass, field
import uuid
@dataclass
class TTSRequest:
"""TTS请求对象(带唯一标识)"""
tts_text: str
mode: str = "预训练音色"
sft_spk: str = ""
seed: int = field(default_factory=lambda: np.random.randint(1, 100000000))
stream: bool = True
speed: float = 1.0
prompt_wav: str = ""
instruct_text: str = ""
request_id: str = field(default_factory=lambda: str(uuid.uuid4())) # 唯一请求ID
class CosyVoiceTTSSocketClient:
"""CosyVoice TTS WebSocket 客户端(异步/流式/带任务队列)"""
def __init__(self, ws_url: str = "ws://localhost:50000/ws/tts", max_queue_size: int = 100):
"""
初始化客户端
:param ws_url: WebSocket 服务端地址
:param max_queue_size: 最大队列长度(防止内存溢出)
"""
self.ws_url = ws_url
self.websocket: Optional[websockets.WebSocketClientProtocol] = None
self.is_connected = False
self.is_processing = False # 是否正在处理请求
# 异步任务队列(FIFO
self.request_queue: asyncio.Queue[TTSRequest] = asyncio.Queue(maxsize=max_queue_size)
# 回调函数定义(所有回调都带request_id,方便关联请求)
self.on_task_enqueue: Callable[[str], None] = lambda req_id: None # 任务入队回调
self.on_start: Callable[[str, Dict[str, Any]], None] = lambda req_id, data: None # 合成开始回调
self.on_audio_chunk: Callable[[str, np.ndarray], None] = lambda req_id, chunk: None # 音频块回调
self.on_end: Callable[[str, Dict[str, Any]], None] = lambda req_id, data: None # 合成结束回调
self.on_error: Callable[[str, str], None] = lambda req_id, msg: None # 错误回调
self.on_queue_full: Callable[[str], None] = lambda req_id: None # 队列满回调
# 音频播放流
self.play_stream: Optional[sd.OutputStream] = None
self.current_req_id: Optional[str] = None
# 是否启用播放功能
self.enable_playback: bool = True
# 外部回调函数,用于返回完整结果
self.external_callback: Optional[Callable[[str, Dict[str, Any]], None]] = None
# 存储每个请求的音频数据
self.audio_buffers: Dict[str, List[np.ndarray]] = {}
async def connect(self):
"""建立 WebSocket 连接(初始化一次)"""
if not self.is_connected:
try:
self.websocket = await websockets.connect(self.ws_url)
self.is_connected = True
print(f"成功连接到 TTS 服务端: {self.ws_url}")
# 启动队列消费协程(后台运行)
asyncio.create_task(self._consume_queue())
except Exception as e:
raise ConnectionError(f"连接失败: {str(e)}")
async def disconnect(self):
"""关闭 WebSocket 连接"""
if self.is_connected and self.websocket:
await self.websocket.close()
self.is_connected = False
self.websocket = None
print("已断开与 TTS 服务端的连接")
if self.play_stream:
self.play_stream.stop()
self.play_stream.close()
self.play_stream = None
def set_external_callback(self, callback: Callable[[str, Dict[str, Any]], None]):
"""设置外部回调函数,用于返回完整结果"""
self.external_callback = callback
def set_playback_enabled(self, enabled: bool):
"""设置是否启用音频播放"""
self.enable_playback = enabled
async def add_tts_request(self, tts_text: str, **kwargs) -> str:
"""
添加TTS请求到队列(异步非阻塞)
:param tts_text: 要合成的文本
:param kwargs: 其他TTS参数
:return: 唯一请求ID
"""
# 创建请求对象
request = TTSRequest(tts_text=tts_text, **kwargs)
req_id = request.request_id
# 初始化音频缓冲区
self.audio_buffers[req_id] = []
# 异步入队
try:
await self.request_queue.put(request)
self.on_task_enqueue(req_id)
print(f"请求 {req_id} 已加入队列,当前队列长度: {self.request_queue.qsize()}")
return req_id
except asyncio.QueueFull:
self.on_queue_full(req_id)
raise Exception(f"队列已满,请求 {req_id} 入队失败")
async def _consume_queue(self):
"""消费队列(后台协程,自动处理排队请求)"""
print("队列消费协程已启动")
while True:
try:
# 等待队列中有请求(阻塞)
request = await self.request_queue.get()
req_id = request.request_id
# 标记为处理中
self.is_processing = True
print(f"开始处理请求 {req_id},剩余队列长度: {self.request_queue.qsize()}")
# 处理当前请求
await self._process_single_request(request)
# 标记任务完成
self.request_queue.task_done()
self.is_processing = False
except Exception as e:
print(f"队列消费异常: {str(e)}")
self.is_processing = False
# 短暂等待后继续消费,避免死循环
await asyncio.sleep(0.1)
async def _process_single_request(self, request: TTSRequest):
"""处理单个TTS请求"""
req_id = request.request_id
# 参数校验
if not request.tts_text:
self.on_error(req_id, "合成文本不能为空")
return
# 确保已连接
if not self.is_connected:
try:
await self.connect()
except Exception as e:
self.on_error(req_id, f"连接失败: {str(e)}")
return
# 构造请求数据
request_data = {
"tts_text": request.tts_text,
"mode": request.mode,
"sft_spk": request.sft_spk,
"seed": request.seed,
"stream": request.stream,
"speed": request.speed,
"prompt_wav": request.prompt_wav,
"instruct_text": request.instruct_text
}
try:
# 发送请求
await self.websocket.send(json.dumps(request_data))
# 处理返回的流式数据
await self._handle_response(req_id)
except websockets.exceptions.ConnectionClosed:
self.is_connected = False
self.on_error(req_id, "连接已关闭")
except Exception as e:
self.on_error(req_id, f"处理请求失败: {str(e)}")
async def _handle_response(self, req_id: str):
"""处理单个请求的服务端响应"""
if not self.websocket:
return
try:
sample_rate = None
while True:
# 接收服务端消息(异步)
response = await self.websocket.recv()
data = json.loads(response)
# 根据状态分发到不同回调(都带req_id)
status = data.get("status")
if status == "start":
# 合成开始 - 返回采样率等信息
sample_rate = data.get("sample_rate")
self.on_start(req_id, data)
elif status == "stream":
# 流式音频块 - 转为 numpy 数组
audio_chunk = np.array(data["audio_chunk"], dtype=np.float32)
# 保存到缓冲区
self.audio_buffers[req_id].append(audio_chunk)
# 音频块回调
self.on_audio_chunk(req_id, audio_chunk)
elif status == "end":
# 合成结束
self.on_end(req_id, data)
# 组装完整音频数据
full_audio = np.concatenate(self.audio_buffers[req_id]) if self.audio_buffers[req_id] else np.array(
[], dtype=np.float32)
# 准备返回给外部回调的数据
result_data = {
"status": "completed",
"request_id": req_id,
"sample_rate": sample_rate,
"audio_data": full_audio,
"audio_length": len(full_audio),
"message": data.get("msg", "合成完成")
}
# 调用外部回调(非阻塞)
if self.external_callback:
# 使用create_task确保回调不会阻塞主流程
asyncio.create_task(self._call_external_callback(req_id, result_data))
# 清理缓冲区
if req_id in self.audio_buffers:
del self.audio_buffers[req_id]
break
elif status == "error":
# 错误处理
error_msg = data.get("msg", "未知错误")
self.on_error(req_id, error_msg)
# 准备错误结果数据
error_data = {
"status": "error",
"request_id": req_id,
"message": error_msg
}
# 调用外部回调
if self.external_callback:
asyncio.create_task(self._call_external_callback(req_id, error_data))
# 清理缓冲区
if req_id in self.audio_buffers:
del self.audio_buffers[req_id]
break
except Exception as e:
error_msg = f"接收数据异常: {str(e)}"
self.on_error(req_id, error_msg)
# 准备错误结果数据
error_data = {
"status": "error",
"request_id": req_id,
"message": error_msg
}
# 调用外部回调
if self.external_callback:
asyncio.create_task(self._call_external_callback(req_id, error_data))
# 清理缓冲区
if req_id in self.audio_buffers:
del self.audio_buffers[req_id]
async def _call_external_callback(self, req_id: str, result_data: Dict[str, Any]):
"""异步调用外部回调函数"""
try:
if self.external_callback:
# 确保回调是异步的,不会阻塞
if asyncio.iscoroutinefunction(self.external_callback):
await self.external_callback(req_id, result_data)
else:
# 如果是同步函数,在线程池中执行
await asyncio.get_event_loop().run_in_executor(
None, self.external_callback, req_id, result_data
)
except Exception as e:
print(f"外部回调执行异常: {str(e)}")
def _play_audio_chunk(self, chunk: np.ndarray):
"""播放音频块(非阻塞)"""
if not self.enable_playback or not self.play_stream:
return
try:
if chunk.size > 0:
self.play_stream.write(chunk)
except Exception as e:
print(f"音频播放异常: {str(e)}")
async def wait_queue_empty(self):
"""等待队列所有任务处理完成(阻塞)"""
await self.request_queue.join()
print("所有队列任务已处理完成")
# ------------------------------
# 使用示例
# ------------------------------
class TTSManager:
"""TTS管理器 - 供外部代码调用"""
def __init__(self, ws_url: str = "ws://localhost:50000/ws/tts"):
self.client = CosyVoiceTTSSocketClient(ws_url)
self._setup_callbacks()
def _setup_callbacks(self):
"""设置内部回调"""
# 合成开始回调
def on_tts_start(req_id, data):
print(f"\n🎤 开始合成 [{req_id[:8]}] - 采样率: {data.get('sample_rate')}")
# 初始化播放流(如果需要播放)
if self.client.enable_playback:
sample_rate = data.get("sample_rate", 24000)
if self.client.play_stream:
self.client.play_stream.stop()
self.client.play_stream.close()
self.client.play_stream = sd.OutputStream(
samplerate=sample_rate,
channels=1,
dtype=np.float32
)
self.client.play_stream.start()
self.client.current_req_id = req_id
# 音频块回调
def on_audio_chunk(req_id, chunk):
if req_id == self.client.current_req_id:
self.client._play_audio_chunk(chunk)
# 合成结束回调
def on_tts_end(req_id, data):
print(f"🏁 合成完成 [{req_id[:8]}] - {data.get('msg', '完成')}")
# 错误回调
def on_tts_error(req_id, msg):
print(f"❌ 合成失败 [{req_id[:8]}] - {msg}")
self.client.on_start = on_tts_start
self.client.on_audio_chunk = on_audio_chunk
self.client.on_end = on_tts_end
self.client.on_error = on_tts_error
async def initialize(self):
"""初始化连接"""
await self.client.connect()
async def shutdown(self):
"""关闭连接"""
await self.client.disconnect()
def set_result_callback(self, callback: Callable[[str, Dict[str, Any]], None]):
"""设置结果回调函数"""
self.client.set_external_callback(callback)
def set_playback_enabled(self, enabled: bool):
"""设置是否启用播放"""
self.client.set_playback_enabled(enabled)
async def synthesize(self, text: str, **kwargs) -> str:
"""
异步合成文本(非阻塞)
:param text: 要合成的文本
:param kwargs: 其他参数(mode, sft_spk, speed等)
:return: 请求ID
"""
return await self.client.add_tts_request(text, **kwargs)
async def wait_all_completed(self):
"""等待所有任务完成"""
await self.client.wait_queue_empty()
# ------------------------------
# 外部调用示例
# ------------------------------
async def external_usage_example():
"""外部代码使用示例"""
# 1. 创建TTS管理器
tts_manager = TTSManager(ws_url="ws://10.10.10.202:50000/ws/tts")
# 2. 设置结果回调函数(接收完整结果)
def handle_tts_result1(req_id: str, result: Dict[str, Any]):
"""处理TTS结果回调"""
status = result.get("status")
if status == "completed":
audio_data = result.get("audio_data")
sample_rate = result.get("sample_rate")
audio_length = result.get("audio_length")
print(f"✅ 收到TTS结果 [{req_id[:8]}]: 长度{audio_length}采样点, 采样率{sample_rate}Hz")
# 这里可以保存音频文件或进行其他处理
# 注意:audio_data是完整的numpy数组
elif status == "error":
error_msg = result.get("message")
print(f"❌ TTS处理失败 [{req_id[:8]}]: {error_msg}")
def handle_tts_result(req_id: str, result: Dict[str, Any]):
"""处理TTS结果回调 - 保存为PCM文件"""
print('处理TTS结果回调', result)
status = result.get("status")
print('status')
if status == "completed":
audio_data = result.get("audio_data")
sample_rate = result.get("sample_rate")
if audio_data is not None and len(audio_data) > 0:
# 转换为PCM数据
pcm_data = (audio_data.astype(np.float32) * 32767).astype(np.int16)
pcm_bytes = pcm_data.tobytes()
# 保存为PCM文件
filename = f"tts_output_{req_id[:8]}.pcm"
with open(filename, "wb") as f:
f.write(pcm_bytes)
print(f"💾 PCM文件已保存: {filename}, 大小: {len(pcm_bytes)}字节")
return filename
elif status == "error":
error_msg = result.get("message")
print(f"❌ TTS处理失败 [{req_id[:8]}]: {error_msg}")
return None
tts_manager.set_result_callback(handle_tts_result)
# 3. 设置是否播放(可选,默认True)
tts_manager.set_playback_enabled(True) # 设置为False则不播放
# 4. 初始化连接
await tts_manager.initialize()
# 5. 异步合成多个文本(非阻塞)
texts = [
"你好,这是第一个排队的TTS请求。",
"我是第二个请求,会等第一个处理完再执行。",
"第三个请求,支持流式播放和队列管理。",
"第四个请求,测试队列的自动消费功能。",
"最后一个请求,处理完成后会自动结束。"
]
req_ids = []
for i, text in enumerate(texts):
req_id = await tts_manager.synthesize(
text,
mode="预训练音色",
sft_spk="中文女",
speed=1.0
)
req_ids.append(req_id)
print(f"已提交请求 {i + 1}: ID={req_id[:8]}")
# 可以立即继续其他操作,不需要等待
await asyncio.sleep(0.5) # 模拟其他操作
# 6. 可以在这里做其他事情,TTS会在后台处理
# 7. 等待所有TTS任务完成(可选)
await tts_manager.wait_all_completed()
# 8. 关闭连接
await tts_manager.shutdown()
async def advanced_usage_example():
"""高级使用示例:动态控制播放和结果处理"""
tts_manager = TTSManager(ws_url="ws://10.10.10.202:50000/ws/tts")
# 自定义结果处理器
class ResultProcessor:
def __init__(self):
self.results = {}
async def process_result(self, req_id: str, result: Dict[str, Any]):
"""异步处理结果"""
if result.get("status") == "completed":
# 保存音频数据等处理
self.results[req_id] = result
print(f"处理完成: {req_id[:8]}, 音频长度: {result.get('audio_length')}")
# 这里可以添加自定义逻辑,如保存到文件、发送到网络等
processor = ResultProcessor()
tts_manager.set_result_callback(processor.process_result)
# 禁用播放,只获取数据
tts_manager.set_playback_enabled(False)
await tts_manager.initialize()
# 批量合成
tasks = []
for i in range(3):
text = f"这是第{i+1}个测试文本,用于演示高级功能。"
task = tts_manager.synthesize(text)
tasks.append(task)
# 并行提交所有请求
req_ids = await asyncio.gather(*tasks)
# 等待完成
await tts_manager.wait_all_completed()
await tts_manager.shutdown()
print(f"所有任务完成,共处理 {len(processor.results)} 个结果")
if __name__ == "__main__":
# 运行示例
try:
# 基础使用示例
asyncio.run(external_usage_example())
# 高级使用示例(取消注释运行)
# asyncio.run(advanced_usage_example())
except KeyboardInterrupt:
print("\n程序被用户中断")
except Exception as e:
print(f"程序异常: {str(e)}")