diff --git a/api_server.py b/api_server.py index 86f2a6d..6141bc4 100644 --- a/api_server.py +++ b/api_server.py @@ -220,23 +220,45 @@ async def create_completion(request: Request, body: CompletionRequest): raise HTTPException(status_code=500, detail=str(e)) @app.post("/v1/audio/speech") -async def create_speech(request: Request): +async def create_speech( + request: Request, +): try: request_data = await request.json() - backend = await _dispatcher.get_backend(request_data.get("model")) + model = request_data.get("model") + voice = request_data.get("voice", "alloy") + response_format = request_data.get("response_format", "mp3") + + if not model: + raise HTTPException( + status_code=400, + detail="Model is required for speech generation" + ) + + backend = await _dispatcher.get_backend(model) if not backend: - raise HTTPException(status_code=400, detail="Unsupported model") + raise HTTPException( + status_code=400, + detail=f"Unsupported model: {model}" + ) + + input_text = request_data.get("input") + if not input_text: + raise HTTPException( + status_code=400, + detail="Input text is required for speech generation" + ) audio_stream = backend.generate_speech( - input_text=request_data.get("input"), - voice=request_data.get("voice", "alloy"), - format=request_data.get("response_format", "mp3") + input_text=input_text, + voice=voice, + format=response_format ) return StreamingResponse( audio_stream, - media_type=f"audio/{request_data.get('response_format', 'mp3')}", - headers={"Content-Disposition": f'attachment; filename="speech.{request_data.get("response_format", "mp3")}"'} + media_type=f"audio/{response_format}", + headers={"Content-Disposition": f'attachment; filename="speech.{response_format}"'} ) except Exception as e: logger.error(f"Speech generation error: {str(e)}") diff --git a/backend/asr_client_backend.py b/backend/asr_client_backend.py index dc979bd..d28b83b 100644 --- a/backend/asr_client_backend.py +++ b/backend/asr_client_backend.py @@ -39,7 +39,6 @@ class ASRClientBackend(BaseModelBackend): for task in self._active_tasks: task.cancel() - self._pool_lock.release() await asyncio.sleep(retry_interval) await asyncio.wait_for(self._pool_lock.acquire(), timeout=timeout - (time.time() - start_time)) diff --git a/backend/tts_client_backend.py b/backend/tts_client_backend.py index e1f33d8..fa4883b 100644 --- a/backend/tts_client_backend.py +++ b/backend/tts_client_backend.py @@ -1,16 +1,18 @@ -from concurrent.futures import ThreadPoolExecutor -from .base_model_backend import BaseModelBackend -from client.tts_client import TTSClient +import time import asyncio +import weakref import base64 import logging import io from pydub import AudioSegment -import numpy as np from typing import AsyncGenerator +from .base_model_backend import BaseModelBackend +from client.tts_client import TTSClient +from concurrent.futures import ThreadPoolExecutor +from services.memory_check import MemoryChecker + class TtsClientBackend(BaseModelBackend): - POOL_SIZE = 1 SUPPORTED_FORMATS = ["mp3", "opus", "aac", "flac", "wav", "pcm"] def __init__(self, model_config): @@ -19,40 +21,62 @@ class TtsClientBackend(BaseModelBackend): self._active_clients = {} self._pool_lock = asyncio.Lock() self.logger = logging.getLogger("api.tts") - self._executor = ThreadPoolExecutor(max_workers=self.POOL_SIZE) + self.POOL_SIZE = 1 + self._inference_executor = ThreadPoolExecutor(max_workers=self.POOL_SIZE) + self._active_tasks = weakref.WeakSet() + self.memory_checker = MemoryChecker( + host=self.config["host"], + port=self.config["port"] + ) self.sample_rate = 16000 self.channels = 1 async def _get_client(self): - async with self._pool_lock: - if self._client_pool: - return self._client_pool.pop() + try: + await asyncio.wait_for(self._pool_lock.acquire(), timeout=30.0) - if len(self._active_clients) >= self.POOL_SIZE: - raise RuntimeError("TTS connection pool exhausted") + start_time = time.time() + timeout = 30.0 + retry_interval = 3 - client = TTSClient( - host=self.config["host"], - port=self.config["port"] - ) - - loop = asyncio.get_event_loop() - await loop.run_in_executor( - self._executor, - lambda: client.setup( - "melotts.setup", - { - "model": self.config["model_name"], - "response_format": "pcm.stream.base64", - "input": "tts.utf-8", - "enoutput": True, - "voice": "alloy" - } + while True: + if self._client_pool: + client = self._client_pool.pop() + return client + + for task in self._active_tasks: + task.cancel() + + self._pool_lock.release() + await asyncio.sleep(retry_interval) + await asyncio.wait_for(self._pool_lock.acquire(), timeout=timeout - (time.time() - start_time)) + + client = TTSClient( + host=self.config["host"], + port=self.config["port"] ) - ) - - self._active_clients[id(client)] = client - return client + self._active_clients[id(client)] = client + + loop = asyncio.get_event_loop() + await loop.run_in_executor( + self._inference_executor, + lambda: client.setup( + "melotts.setup", + { + "model": self.config["model_name"], + "response_format": "pcm.stream.base64", + "input": "tts.utf-8", + "enoutput": True, + "voice": "alloy" + } + ) + ) + return client + except asyncio.TimeoutError: + raise RuntimeError("Server busy, please try again later.") + finally: + if self._pool_lock.locked(): + self._pool_lock.release() async def _release_client(self, client): async with self._pool_lock: @@ -97,25 +121,15 @@ class TtsClientBackend(BaseModelBackend): async def generate_speech(self, input_text: str, voice: str = "alloy", format: str = "mp3") -> AsyncGenerator[bytes, None]: client = await self._get_client() + task = asyncio.current_task() + self._active_tasks.add(task) + full_data = b'' try: loop = asyncio.get_event_loop() - sync_gen = client.inference_stream(input_text, object_type="tts.utf-8") - - def safe_next(): - try: - return next(sync_gen) - except StopIteration: - return None - - full_data = b'' - while True: - chunk = await loop.run_in_executor(self._executor, safe_next) - if chunk is None: - break - + async for chunk in client.inference_stream(input_text, object_type="tts.utf-8"): pcm_data = base64.b64decode(chunk) encoded_data = await loop.run_in_executor( - self._executor, + self._inference_executor, self._encode_audio, pcm_data, format @@ -130,4 +144,5 @@ class TtsClientBackend(BaseModelBackend): yield final_audio finally: - await self._release_client(client) \ No newline at end of file + self._active_tasks.discard(task) + await self._release_client(client) \ No newline at end of file diff --git a/client/tts_client.py b/client/tts_client.py index 662a596..162e34e 100644 --- a/client/tts_client.py +++ b/client/tts_client.py @@ -2,10 +2,11 @@ import json import socket import time import uuid -from typing import Generator +from typing import Generator, AsyncGenerator import logging import threading import base64 +import asyncio logger = logging.getLogger("tts_client") logger.setLevel(logging.DEBUG) @@ -65,18 +66,20 @@ class TTSClient: request_id = self._send_request("setup", object, model_config) return self._wait_response(request_id) - def inference_stream(self, query: str, object_type: str = "llm.utf-8") -> Generator[str, None, None]: + async def inference_stream(self, query: str, object_type: str = "llm.utf-8") -> AsyncGenerator[str, None]: request_id = self._send_request("inference", object_type, query) buffer = b'' - + + loop = asyncio.get_event_loop() + while True: start_time = time.time() while time.time() - start_time < 3600: - chunk = self.sock.recv(4096) + chunk = await loop.run_in_executor(None, self.sock.recv, 4096) if not chunk: break buffer += chunk - + while b'\n' in buffer: line, buffer = buffer.split(b'\n', 1) try: