server.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322
  1. #!/usr/bin/env python3
  2. import asyncio
  3. import base64
  4. import json
  5. import argparse
  6. import signal
  7. import sys
  8. import os
  9. import io
  10. import wave
  11. import hashlib
  12. from concurrent.futures import ThreadPoolExecutor
  13. from functools import lru_cache
  14. from typing import Set, Optional, Tuple, List, Dict
  15. import numpy as np
  16. import websockets
  17. from websockets.server import WebSocketServerProtocol
  18. sys.stdout = os.fdopen(sys.stdout.fileno(), 'w', buffering=1)
  19. class AudioCache:
  20. def __init__(self, max_size: int = 200):
  21. self.cache: Dict[str, bytes] = {}
  22. self.access_order: List[str] = []
  23. self.max_size = max_size
  24. self.hits = 0
  25. self.misses = 0
  26. def _make_key(self, text: str, ref_hash: str) -> str:
  27. return hashlib.md5(f"{ref_hash}:{text}".encode()).hexdigest()
  28. def get(self, text: str, ref_hash: str) -> Optional[bytes]:
  29. key = self._make_key(text, ref_hash)
  30. if key in self.cache:
  31. self.access_order.remove(key)
  32. self.access_order.append(key)
  33. self.hits += 1
  34. return self.cache[key]
  35. self.misses += 1
  36. return None
  37. def put(self, text: str, ref_hash: str, audio_data: bytes):
  38. key = self._make_key(text, ref_hash)
  39. if key in self.cache:
  40. self.access_order.remove(key)
  41. elif len(self.cache) >= self.max_size:
  42. oldest = self.access_order.pop(0)
  43. del self.cache[oldest]
  44. self.cache[key] = audio_data
  45. self.access_order.append(key)
  46. def get_stats(self) -> dict:
  47. total = self.hits + self.misses
  48. hit_rate = (self.hits / total * 100) if total > 0 else 0
  49. return {
  50. 'size': len(self.cache),
  51. 'max_size': self.max_size,
  52. 'hits': self.hits,
  53. 'misses': self.misses,
  54. 'hit_rate': f"{hit_rate:.1f}%"
  55. }
  56. class TTSEngine:
  57. def __init__(self, model_id: str = "Qwen/Qwen3-TTS-12Hz-0.6B-Base",
  58. ref_audio_path: str = None):
  59. import soundfile as sf
  60. self.model_id = model_id
  61. self.sample_rate = 24000
  62. self.ref_audio_path = ref_audio_path
  63. self.ref_hash = None
  64. self.audio_cache = AudioCache(max_size=200)
  65. print(f"Loading Qwen3-TTS model: {model_id}")
  66. from qwen_tts import Qwen3TTSModel
  67. self.model = Qwen3TTSModel.from_pretrained(model_id, device_map='mps')
  68. print("Model loaded successfully")
  69. if self.ref_audio_path and os.path.exists(self.ref_audio_path):
  70. self.ref_audio, self.ref_sr = self._load_audio(self.ref_audio_path)
  71. self.ref_hash = hashlib.md5(open(self.ref_audio_path, 'rb').read()).hexdigest()
  72. self.prompt_items = self.model.create_voice_clone_prompt(
  73. ref_audio=(self.ref_audio, self.ref_sr),
  74. ref_text="",
  75. x_vector_only_mode=True
  76. )
  77. print(f"Reference audio loaded: {self.ref_audio_path}")
  78. else:
  79. self.ref_audio = None
  80. self.ref_sr = None
  81. self.prompt_items = None
  82. print("Warning: No reference audio provided")
  83. def _load_audio(self, path: str) -> Tuple[np.ndarray, int]:
  84. import soundfile as sf
  85. audio, sr = sf.read(path)
  86. if len(audio.shape) > 1:
  87. audio = audio[:, 0]
  88. return audio, sr
  89. def synthesize(self, text: str, ref_audio_path: str = None) -> bytes:
  90. if ref_audio_path:
  91. ref_audio, ref_sr = self._load_audio(ref_audio_path)
  92. ref_hash = hashlib.md5(open(ref_audio_path, 'rb').read()).hexdigest()
  93. prompt = self.model.create_voice_clone_prompt(
  94. ref_audio=(ref_audio, ref_sr),
  95. ref_text="",
  96. x_vector_only_mode=True
  97. )
  98. elif self.ref_hash:
  99. ref_hash = self.ref_hash
  100. prompt = self.prompt_items
  101. else:
  102. raise ValueError("Reference audio required for Base model")
  103. cached = self.audio_cache.get(text, ref_hash)
  104. if cached:
  105. return cached
  106. audio_chunks, sample_rate = self.model.generate_voice_clone(
  107. text=text,
  108. voice_clone_prompt=prompt,
  109. x_vector_only_mode=True,
  110. )
  111. wav_data = self._audio_to_wav(audio_chunks, sample_rate)
  112. self.audio_cache.put(text, ref_hash, wav_data)
  113. return wav_data
  114. def _audio_to_wav(self, audio_data, sample_rate: int = None) -> bytes:
  115. if sample_rate is None:
  116. sample_rate = self.sample_rate
  117. if isinstance(audio_data, list):
  118. audio_array = np.concatenate(audio_data) if len(audio_data) > 0 else np.array([])
  119. else:
  120. audio_array = audio_data
  121. if audio_array.dtype != np.int16:
  122. audio_array = (audio_array * 32767).astype(np.int16)
  123. wav_io = io.BytesIO()
  124. with wave.open(wav_io, 'wb') as wav_file:
  125. wav_file.setnchannels(1)
  126. wav_file.setsampwidth(2)
  127. wav_file.setframerate(sample_rate)
  128. wav_file.writeframes(audio_array.tobytes())
  129. return wav_io.getvalue()
  130. def get_model_info(self) -> dict:
  131. info = {
  132. 'model_id': self.model_id,
  133. 'sample_rate': self.sample_rate,
  134. 'has_ref_audio': self.ref_audio is not None,
  135. 'cache_stats': self.audio_cache.get_stats()
  136. }
  137. return info
  138. class TTSServer:
  139. def __init__(
  140. self,
  141. model_id: str = "Qwen/Qwen3-TTS-12Hz-0.6B-Base",
  142. port: int = 8766,
  143. host: str = "localhost",
  144. ref_audio_path: str = None,
  145. max_workers: int = 4
  146. ):
  147. self.port = port
  148. self.host = host
  149. self.ref_audio_path = ref_audio_path
  150. self.clients: Set[WebSocketServerProtocol] = set()
  151. self.is_running = True
  152. self.executor = ThreadPoolExecutor(max_workers=max_workers)
  153. print(f"Initializing TTS engine...")
  154. self.engine = TTSEngine(model_id, ref_audio_path)
  155. print(f"TTS Engine ready (max_workers={max_workers})")
  156. async def register(self, websocket: WebSocketServerProtocol):
  157. self.clients.add(websocket)
  158. print(f"Client connected: {websocket.remote_address}")
  159. await self.send_message(websocket, {
  160. "type": "connected",
  161. "model": self.engine.get_model_info()
  162. })
  163. async def unregister(self, websocket: WebSocketServerProtocol):
  164. self.clients.discard(websocket)
  165. print(f"Client disconnected: {websocket.remote_address}")
  166. async def send_message(self, websocket: WebSocketServerProtocol, message: dict):
  167. if websocket in self.clients:
  168. try:
  169. await websocket.send(json.dumps(message, ensure_ascii=False))
  170. except websockets.exceptions.ConnectionClosed:
  171. await self.unregister(websocket)
  172. async def handle_synthesize(self, websocket: WebSocketServerProtocol,
  173. text: str, ref_audio: str = None, chunk_index: int = -1):
  174. loop = asyncio.get_event_loop()
  175. try:
  176. print(f"Starting synthesis for chunk {chunk_index}: {text[:30]}...")
  177. audio_data = await loop.run_in_executor(
  178. self.executor,
  179. self.engine.synthesize,
  180. text,
  181. ref_audio
  182. )
  183. audio_base64 = base64.b64encode(audio_data).decode('utf-8')
  184. await self.send_message(websocket, {
  185. "type": "audio",
  186. "data": audio_base64,
  187. "format": "wav",
  188. "sample_rate": self.engine.sample_rate,
  189. "chunk_index": chunk_index,
  190. "is_first": True,
  191. "is_last": True
  192. })
  193. print(f"Sent audio for chunk {chunk_index} ({len(audio_data)} bytes)")
  194. except Exception as e:
  195. print(f"Synthesis error for chunk {chunk_index}: {e}")
  196. await self.send_message(websocket, {
  197. "type": "error",
  198. "message": str(e),
  199. "chunk_index": chunk_index
  200. })
  201. async def handle_message(self, websocket: WebSocketServerProtocol, message: dict):
  202. msg_type = message.get("type")
  203. if msg_type == "synthesize":
  204. text = message.get("text", "")
  205. ref_audio = message.get("ref_audio")
  206. chunk_index = message.get("chunk_index", -1)
  207. if not text:
  208. await self.send_message(websocket, {
  209. "type": "error",
  210. "message": "No text provided"
  211. })
  212. return
  213. asyncio.create_task(self.handle_synthesize(websocket, text, ref_audio, chunk_index))
  214. elif msg_type == "voices":
  215. await self.send_message(websocket, {
  216. "type": "voices",
  217. "voices": ["default"]
  218. })
  219. elif msg_type == "cache_stats":
  220. stats = self.engine.audio_cache.get_stats()
  221. await self.send_message(websocket, {
  222. "type": "cache_stats",
  223. **stats
  224. })
  225. async def handler(self, websocket: WebSocketServerProtocol):
  226. await self.register(websocket)
  227. try:
  228. async for raw_message in websocket:
  229. try:
  230. if isinstance(raw_message, str):
  231. message = json.loads(raw_message)
  232. await self.handle_message(websocket, message)
  233. except json.JSONDecodeError:
  234. print(f"Invalid JSON from {websocket.remote_address}")
  235. except Exception as e:
  236. print(f"Error handling message: {e}")
  237. except websockets.exceptions.ConnectionClosed:
  238. pass
  239. finally:
  240. await self.unregister(websocket)
  241. async def start(self):
  242. print(f"Starting TTS server on {self.host}:{self.port}")
  243. async with websockets.serve(self.handler, self.host, self.port):
  244. await asyncio.Future()
  245. def run(self):
  246. asyncio.run(self.start())
  247. def main():
  248. parser = argparse.ArgumentParser(description="TTS WebSocket Server (Optimized with Cache)")
  249. parser.add_argument("--model", "-m", default="Qwen/Qwen3-TTS-12Hz-0.6B-Base")
  250. parser.add_argument("--port", "-p", type=int, default=8766)
  251. parser.add_argument("--host", default="localhost")
  252. parser.add_argument("--ref-audio", "-r", default=None)
  253. parser.add_argument("--workers", "-w", type=int, default=4)
  254. args = parser.parse_args()
  255. server = TTSServer(
  256. model_id=args.model,
  257. port=args.port,
  258. host=args.host,
  259. ref_audio_path=args.ref_audio,
  260. max_workers=args.workers
  261. )
  262. def signal_handler(sig, frame):
  263. print("\nShutting down server...")
  264. server.is_running = False
  265. sys.exit(0)
  266. signal.signal(signal.SIGINT, signal_handler)
  267. signal.signal(signal.SIGTERM, signal_handler)
  268. server.run()
  269. if __name__ == "__main__":
  270. main()