[refactor] Refactor tts_client_backend

This commit is contained in:
LittleMouse
2025-04-14 14:16:27 +08:00
parent 46e0ecf3b2
commit d54df28357
4 changed files with 101 additions and 62 deletions
+30 -8
View File
@@ -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)}")
-1
View File
@@ -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))
+63 -48
View File
@@ -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)
self._active_tasks.discard(task)
await self._release_client(client)
+8 -5
View File
@@ -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: