mirror of
https://github.com/m5stack/ModuleLLM-OpenAI-Plugin.git
synced 2026-05-20 11:37:26 -07:00
[refactor] Refactor tts_client_backend
This commit is contained in:
+30
-8
@@ -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)}")
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user