mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-09-19 23:00:45 -07:00
Implement TTS provider management system and update release workflow
- Added support for TTS providers in the backend, including endpoints for listing, starting, stopping, and downloading providers. - Enhanced the release workflow to build and upload TTS provider binaries for both Windows and Linux platforms. - Updated the architecture documentation to reflect the new provider system and its benefits for modularity and user experience. - Introduced a new `ProviderSettings` component in the frontend for managing provider configurations.
This commit is contained in:
+15
-17
@@ -30,7 +30,7 @@ def build_server():
|
||||
args.extend(['--paths', str(qwen_tts_path)])
|
||||
print(f"Using local qwen_tts source from: {qwen_tts_path}")
|
||||
|
||||
# Add common hidden imports
|
||||
# Add common hidden imports (always included)
|
||||
args.extend([
|
||||
'--hidden-import', 'backend',
|
||||
'--hidden-import', 'backend.main',
|
||||
@@ -42,38 +42,30 @@ def build_server():
|
||||
'--hidden-import', 'backend.tts',
|
||||
'--hidden-import', 'backend.transcribe',
|
||||
'--hidden-import', 'backend.platform_detect',
|
||||
'--hidden-import', 'backend.backends',
|
||||
'--hidden-import', 'backend.backends.pytorch_backend',
|
||||
'--hidden-import', 'backend.providers',
|
||||
'--hidden-import', 'backend.providers.base',
|
||||
'--hidden-import', 'backend.providers.bundled',
|
||||
'--hidden-import', 'backend.providers.types',
|
||||
'--hidden-import', 'backend.utils.audio',
|
||||
'--hidden-import', 'backend.utils.cache',
|
||||
'--hidden-import', 'backend.utils.progress',
|
||||
'--hidden-import', 'backend.utils.hf_progress',
|
||||
'--hidden-import', 'backend.utils.validation',
|
||||
'--hidden-import', 'torch',
|
||||
'--hidden-import', 'transformers',
|
||||
'--hidden-import', 'fastapi',
|
||||
'--hidden-import', 'uvicorn',
|
||||
'--hidden-import', 'sqlalchemy',
|
||||
'--hidden-import', 'librosa',
|
||||
'--hidden-import', 'soundfile',
|
||||
'--hidden-import', 'qwen_tts',
|
||||
'--hidden-import', 'qwen_tts.inference',
|
||||
'--hidden-import', 'qwen_tts.inference.qwen3_tts_model',
|
||||
'--hidden-import', 'qwen_tts.inference.qwen3_tts_tokenizer',
|
||||
'--hidden-import', 'qwen_tts.core',
|
||||
'--hidden-import', 'qwen_tts.cli',
|
||||
'--copy-metadata', 'qwen-tts',
|
||||
'--collect-submodules', 'qwen_tts',
|
||||
'--collect-data', 'qwen_tts',
|
||||
# Fix for pkg_resources and jaraco namespace packages
|
||||
'--hidden-import', 'pkg_resources.extern',
|
||||
'--collect-submodules', 'jaraco',
|
||||
])
|
||||
|
||||
# Add MLX-specific imports if building on Apple Silicon
|
||||
# Platform-specific TTS backend handling
|
||||
if is_apple_silicon():
|
||||
print("Building for Apple Silicon - including MLX dependencies")
|
||||
print("Building for Apple Silicon - including MLX dependencies (bundled)")
|
||||
args.extend([
|
||||
'--hidden-import', 'backend.backends',
|
||||
'--hidden-import', 'backend.backends.mlx_backend',
|
||||
'--hidden-import', 'mlx',
|
||||
'--hidden-import', 'mlx.core',
|
||||
@@ -88,7 +80,13 @@ def build_server():
|
||||
'--collect-data', 'mlx_audio',
|
||||
])
|
||||
else:
|
||||
print("Building for non-Apple Silicon platform - PyTorch only")
|
||||
print("Building for Windows/Linux - excluding PyTorch/Qwen-TTS (providers downloaded separately)")
|
||||
# Note: PyTorch and Qwen-TTS are NOT included - users will download providers separately
|
||||
# Only include backend abstraction (no actual TTS implementation)
|
||||
args.extend([
|
||||
'--hidden-import', 'backend.backends',
|
||||
'--hidden-import', 'backend.backends.pytorch_backend', # Keep for reference, but won't work without PyTorch
|
||||
])
|
||||
|
||||
args.extend([
|
||||
'--noconfirm',
|
||||
|
||||
+200
-15
@@ -29,6 +29,8 @@ from .utils.progress import get_progress_manager
|
||||
from .utils.tasks import get_task_manager
|
||||
from .utils.cache import clear_voice_prompt_cache
|
||||
from .platform_detect import get_backend_type
|
||||
from .providers import get_provider_manager
|
||||
from .providers.types import ProviderType
|
||||
|
||||
app = FastAPI(
|
||||
title="voicebox API",
|
||||
@@ -74,7 +76,7 @@ async def health():
|
||||
from pathlib import Path
|
||||
import os
|
||||
|
||||
tts_model = tts.get_tts_model()
|
||||
tts_model = await tts.get_tts_model_async()
|
||||
backend_type = get_backend_type()
|
||||
|
||||
# Check for GPU availability (CUDA or MPS)
|
||||
@@ -549,7 +551,7 @@ async def generate_speech(
|
||||
)
|
||||
|
||||
# Generate audio
|
||||
tts_model = tts.get_tts_model()
|
||||
tts_model = await tts.get_tts_model_async()
|
||||
# Load the requested model size if different from current (async to not block)
|
||||
model_size = data.model_size or "1.7B"
|
||||
|
||||
@@ -1113,8 +1115,8 @@ async def get_sample_audio(sample_id: str, db: Session = Depends(get_db)):
|
||||
async def load_model(model_size: str = "1.7B"):
|
||||
"""Manually load TTS model."""
|
||||
try:
|
||||
tts_model = tts.get_tts_model()
|
||||
await tts_model.load_model_async(model_size)
|
||||
tts_model = await tts.get_tts_model_async()
|
||||
await tts_model.load_model(model_size)
|
||||
return {"message": f"Model {model_size} loaded successfully"}
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
@@ -1172,10 +1174,10 @@ async def get_model_status():
|
||||
except ImportError:
|
||||
use_scan_cache = False
|
||||
|
||||
def check_tts_loaded(model_size: str):
|
||||
async def check_tts_loaded(model_size: str):
|
||||
"""Check if TTS model is loaded with specific size."""
|
||||
try:
|
||||
tts_model = tts.get_tts_model()
|
||||
tts_model = await tts.get_tts_model_async()
|
||||
return tts_model.is_loaded() and getattr(tts_model, 'model_size', None) == model_size
|
||||
except Exception:
|
||||
return False
|
||||
@@ -1211,14 +1213,14 @@ async def get_model_status():
|
||||
"display_name": "Qwen TTS 1.7B",
|
||||
"hf_repo_id": tts_1_7b_id,
|
||||
"model_size": "1.7B",
|
||||
"check_loaded": lambda: check_tts_loaded("1.7B"),
|
||||
"check_loaded": lambda: check_tts_loaded("1.7B"), # Async function
|
||||
},
|
||||
{
|
||||
"model_name": "qwen-tts-0.6B",
|
||||
"display_name": "Qwen TTS 0.6B",
|
||||
"hf_repo_id": tts_0_6b_id,
|
||||
"model_size": "0.6B",
|
||||
"check_loaded": lambda: check_tts_loaded("0.6B"),
|
||||
"check_loaded": lambda: check_tts_loaded("0.6B"), # Async function
|
||||
},
|
||||
{
|
||||
"model_name": "whisper-base",
|
||||
@@ -1356,7 +1358,11 @@ async def get_model_status():
|
||||
|
||||
# Check if loaded in memory
|
||||
try:
|
||||
loaded = config["check_loaded"]()
|
||||
check_func = config["check_loaded"]
|
||||
if asyncio.iscoroutinefunction(check_func):
|
||||
loaded = await check_func()
|
||||
else:
|
||||
loaded = check_func()
|
||||
except Exception:
|
||||
loaded = False
|
||||
|
||||
@@ -1379,7 +1385,11 @@ async def get_model_status():
|
||||
except Exception as e:
|
||||
# If check fails, try to at least check if loaded
|
||||
try:
|
||||
loaded = config["check_loaded"]()
|
||||
check_func = config["check_loaded"]
|
||||
if asyncio.iscoroutinefunction(check_func):
|
||||
loaded = await check_func()
|
||||
else:
|
||||
loaded = check_func()
|
||||
except Exception:
|
||||
loaded = False
|
||||
|
||||
@@ -1406,14 +1416,24 @@ async def trigger_model_download(request: models.ModelDownloadRequest):
|
||||
task_manager = get_task_manager()
|
||||
progress_manager = get_progress_manager()
|
||||
|
||||
async def load_tts_model_1_7b():
|
||||
"""Load 1.7B TTS model."""
|
||||
tts_model = await tts.get_tts_model_async()
|
||||
await tts_model.load_model("1.7B")
|
||||
|
||||
async def load_tts_model_0_6b():
|
||||
"""Load 0.6B TTS model."""
|
||||
tts_model = await tts.get_tts_model_async()
|
||||
await tts_model.load_model("0.6B")
|
||||
|
||||
model_configs = {
|
||||
"qwen-tts-1.7B": {
|
||||
"model_size": "1.7B",
|
||||
"load_func": lambda: tts.get_tts_model().load_model("1.7B"),
|
||||
"load_func": load_tts_model_1_7b,
|
||||
},
|
||||
"qwen-tts-0.6B": {
|
||||
"model_size": "0.6B",
|
||||
"load_func": lambda: tts.get_tts_model().load_model("0.6B"),
|
||||
"load_func": load_tts_model_0_6b,
|
||||
},
|
||||
"whisper-base": {
|
||||
"model_size": "base",
|
||||
@@ -1472,6 +1492,171 @@ async def trigger_model_download(request: models.ModelDownloadRequest):
|
||||
return {"message": f"Model {request.model_name} download started"}
|
||||
|
||||
|
||||
# ============================================
|
||||
# PROVIDER ENDPOINTS
|
||||
# ============================================
|
||||
|
||||
@app.get("/providers")
|
||||
async def list_providers():
|
||||
"""List all available provider types."""
|
||||
manager = get_provider_manager()
|
||||
installed = await manager.list_installed()
|
||||
|
||||
# Get info for all known provider types
|
||||
all_providers = [
|
||||
"bundled-mlx",
|
||||
"bundled-pytorch",
|
||||
"pytorch-cpu",
|
||||
"pytorch-cuda",
|
||||
"remote",
|
||||
"openai",
|
||||
]
|
||||
|
||||
providers_info = []
|
||||
for provider_type in all_providers:
|
||||
info = await manager.get_provider_info(provider_type)
|
||||
providers_info.append(info)
|
||||
|
||||
return {
|
||||
"providers": providers_info,
|
||||
"installed": installed,
|
||||
}
|
||||
|
||||
|
||||
@app.get("/providers/installed")
|
||||
async def list_installed_providers():
|
||||
"""List installed provider types."""
|
||||
manager = get_provider_manager()
|
||||
installed = await manager.list_installed()
|
||||
return {"installed": installed}
|
||||
|
||||
|
||||
@app.get("/providers/active")
|
||||
async def get_active_provider():
|
||||
"""Get information about the currently active provider."""
|
||||
manager = get_provider_manager()
|
||||
provider = await manager.get_active_provider()
|
||||
|
||||
health = await provider.health()
|
||||
status = await provider.status()
|
||||
|
||||
return {
|
||||
"provider": health["provider"],
|
||||
"health": health,
|
||||
"status": status,
|
||||
}
|
||||
|
||||
|
||||
@app.post("/providers/start")
|
||||
async def start_provider(data: dict):
|
||||
"""Start a specific provider."""
|
||||
provider_type = data.get("provider_type")
|
||||
if not provider_type:
|
||||
raise HTTPException(status_code=400, detail="provider_type is required")
|
||||
|
||||
manager = get_provider_manager()
|
||||
try:
|
||||
await manager.start_provider(provider_type)
|
||||
provider = await manager.get_active_provider()
|
||||
health = await provider.health()
|
||||
return {
|
||||
"message": f"Provider {provider_type} started",
|
||||
"provider": health,
|
||||
}
|
||||
except NotImplementedError as e:
|
||||
raise HTTPException(status_code=501, detail=str(e))
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@app.post("/providers/stop")
|
||||
async def stop_provider():
|
||||
"""Stop the currently active provider."""
|
||||
manager = get_provider_manager()
|
||||
await manager.stop_provider()
|
||||
return {"message": "Provider stopped"}
|
||||
|
||||
|
||||
@app.post("/providers/download")
|
||||
async def download_provider_endpoint(data: dict):
|
||||
"""Download a provider binary."""
|
||||
from .providers.installer import download_provider
|
||||
|
||||
provider_type = data.get("provider_type")
|
||||
if not provider_type:
|
||||
raise HTTPException(status_code=400, detail="provider_type is required")
|
||||
|
||||
if provider_type not in ["pytorch-cpu", "pytorch-cuda"]:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Provider type {provider_type} cannot be downloaded"
|
||||
)
|
||||
|
||||
try:
|
||||
# Start download in background
|
||||
asyncio.create_task(download_provider(provider_type))
|
||||
return {
|
||||
"message": f"Provider {provider_type} download started",
|
||||
"provider_type": provider_type,
|
||||
}
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@app.get("/providers/download/progress/{provider_type}")
|
||||
async def get_provider_download_progress(provider_type: str):
|
||||
"""Get provider download progress via Server-Sent Events."""
|
||||
from fastapi.responses import StreamingResponse
|
||||
from .utils.progress import get_progress_manager
|
||||
|
||||
progress_manager = get_progress_manager()
|
||||
|
||||
async def event_generator():
|
||||
"""Generate SSE events for provider download progress."""
|
||||
import asyncio
|
||||
import json
|
||||
|
||||
last_progress = None
|
||||
|
||||
while True:
|
||||
progress = progress_manager.get_progress(provider_type)
|
||||
|
||||
if progress and progress != last_progress:
|
||||
yield f"data: {json.dumps(progress)}\n\n"
|
||||
last_progress = progress
|
||||
|
||||
if progress.get("status") in ["complete", "error"]:
|
||||
break
|
||||
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
return StreamingResponse(event_generator(), media_type="text/event-stream")
|
||||
|
||||
|
||||
@app.delete("/providers/{provider_type}")
|
||||
async def delete_provider_endpoint(provider_type: str):
|
||||
"""Delete an installed provider."""
|
||||
from .providers.installer import delete_provider
|
||||
|
||||
if provider_type not in ["pytorch-cpu", "pytorch-cuda"]:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Provider type {provider_type} cannot be deleted"
|
||||
)
|
||||
|
||||
deleted = delete_provider(provider_type)
|
||||
|
||||
if deleted:
|
||||
return {"message": f"Provider {provider_type} deleted successfully"}
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail=f"Provider {provider_type} not found"
|
||||
)
|
||||
|
||||
|
||||
@app.delete("/models/{model_name}")
|
||||
async def delete_model(model_name: str):
|
||||
"""Delete a downloaded model from the HuggingFace cache."""
|
||||
@@ -1522,9 +1707,9 @@ async def delete_model(model_name: str):
|
||||
try:
|
||||
# Check if model is loaded and unload it first
|
||||
if config["model_type"] == "tts":
|
||||
tts_model = tts.get_tts_model()
|
||||
if tts_model.is_loaded() and tts_model.model_size == config["model_size"]:
|
||||
tts.unload_tts_model()
|
||||
tts_model = await tts.get_tts_model_async()
|
||||
if tts_model.is_loaded() and getattr(tts_model, 'model_size', None) == config["model_size"]:
|
||||
tts_model.unload_model()
|
||||
elif config["model_type"] == "whisper":
|
||||
whisper_model = transcribe.get_whisper_model()
|
||||
if whisper_model.is_loaded() and whisper_model.model_size == config["model_size"]:
|
||||
|
||||
@@ -0,0 +1,220 @@
|
||||
"""
|
||||
Provider management system for TTS providers.
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
import platform
|
||||
from pathlib import Path
|
||||
|
||||
from .base import TTSProvider
|
||||
from .types import ProviderType
|
||||
from .bundled import BundledProvider
|
||||
from .local import LocalProvider
|
||||
from .installer import get_provider_binary_path
|
||||
from ..config import get_data_dir
|
||||
import subprocess
|
||||
import socket
|
||||
|
||||
|
||||
class ProviderManager:
|
||||
"""Manages TTS provider lifecycle."""
|
||||
|
||||
def __init__(self):
|
||||
self.active_provider: Optional[TTSProvider] = None
|
||||
self._default_provider: Optional[TTSProvider] = None
|
||||
self._provider_process: Optional[subprocess.Popen] = None
|
||||
self._provider_port: Optional[int] = None
|
||||
|
||||
def _get_default_provider(self) -> TTSProvider:
|
||||
"""Get the default bundled provider."""
|
||||
if self._default_provider is None:
|
||||
self._default_provider = BundledProvider()
|
||||
return self._default_provider
|
||||
|
||||
async def get_active_provider(self) -> TTSProvider:
|
||||
"""
|
||||
Get the currently active provider.
|
||||
|
||||
Returns:
|
||||
Active TTS provider instance
|
||||
"""
|
||||
if self.active_provider is None:
|
||||
# Default to bundled provider
|
||||
self.active_provider = self._get_default_provider()
|
||||
return self.active_provider
|
||||
|
||||
async def start_provider(self, provider_type: str) -> None:
|
||||
"""
|
||||
Start a TTS provider.
|
||||
|
||||
Args:
|
||||
provider_type: Type of provider to start
|
||||
"""
|
||||
if provider_type in ["bundled-mlx", "bundled-pytorch"]:
|
||||
# Use bundled provider
|
||||
self.active_provider = self._get_default_provider()
|
||||
elif provider_type in ["pytorch-cpu", "pytorch-cuda"]:
|
||||
# Start local provider subprocess
|
||||
provider_path = get_provider_binary_path(provider_type)
|
||||
if not provider_path or not provider_path.exists():
|
||||
raise ValueError(f"Provider {provider_type} is not installed. Please download it first.")
|
||||
|
||||
# Find a free port
|
||||
port = self._get_free_port()
|
||||
|
||||
# Start provider subprocess
|
||||
from ..config import get_data_dir
|
||||
process = subprocess.Popen(
|
||||
[
|
||||
str(provider_path),
|
||||
"--port", str(port),
|
||||
"--data-dir", str(get_data_dir()),
|
||||
],
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
)
|
||||
|
||||
# Wait for provider to be ready
|
||||
base_url = f"http://127.0.0.1:{port}"
|
||||
await self._wait_for_provider_health(base_url, timeout=30)
|
||||
|
||||
# Create LocalProvider instance
|
||||
self.active_provider = LocalProvider(base_url)
|
||||
self._provider_process = process
|
||||
self._provider_port = port
|
||||
elif provider_type == "remote":
|
||||
# Remote provider - will be implemented in Phase 5
|
||||
raise NotImplementedError("Remote provider not yet implemented")
|
||||
elif provider_type == "openai":
|
||||
# OpenAI provider - will be implemented in Phase 5
|
||||
raise NotImplementedError("OpenAI provider not yet implemented")
|
||||
else:
|
||||
raise ValueError(f"Unknown provider type: {provider_type}")
|
||||
|
||||
async def stop_provider(self) -> None:
|
||||
"""Stop the active provider."""
|
||||
if self.active_provider:
|
||||
# Only stop if it's not the default bundled provider
|
||||
if self.active_provider is not self._default_provider:
|
||||
if hasattr(self.active_provider, 'stop'):
|
||||
await self.active_provider.stop()
|
||||
self.active_provider = None
|
||||
|
||||
# Stop subprocess if running
|
||||
if self._provider_process:
|
||||
self._provider_process.terminate()
|
||||
try:
|
||||
self._provider_process.wait(timeout=5)
|
||||
except subprocess.TimeoutExpired:
|
||||
self._provider_process.kill()
|
||||
self._provider_process = None
|
||||
self._provider_port = None
|
||||
|
||||
async def list_installed(self) -> list[str]:
|
||||
"""
|
||||
List installed provider types.
|
||||
|
||||
Returns:
|
||||
List of installed provider type strings
|
||||
"""
|
||||
installed = []
|
||||
|
||||
# Bundled providers are always available
|
||||
system = platform.system()
|
||||
machine = platform.machine()
|
||||
|
||||
if system == "Darwin" and machine == "arm64":
|
||||
installed.append("bundled-mlx")
|
||||
else:
|
||||
installed.append("bundled-pytorch")
|
||||
|
||||
# Check for downloaded providers (Phase 2)
|
||||
providers_dir = _get_providers_dir()
|
||||
if providers_dir.exists():
|
||||
for provider_file in providers_dir.glob("tts-provider-*"):
|
||||
if provider_file.is_file() and provider_file.stat().st_size > 0:
|
||||
name = provider_file.name
|
||||
if "pytorch-cpu" in name:
|
||||
installed.append("pytorch-cpu")
|
||||
elif "pytorch-cuda" in name:
|
||||
installed.append("pytorch-cuda")
|
||||
|
||||
return installed
|
||||
|
||||
async def get_provider_info(self, provider_type: str) -> dict:
|
||||
"""
|
||||
Get information about a provider.
|
||||
|
||||
Args:
|
||||
provider_type: Type of provider
|
||||
|
||||
Returns:
|
||||
Provider information dictionary
|
||||
"""
|
||||
if provider_type in ["bundled-mlx", "bundled-pytorch"]:
|
||||
return {
|
||||
"type": provider_type,
|
||||
"name": "Bundled Provider",
|
||||
"installed": True,
|
||||
"size_mb": None, # Bundled, no separate size
|
||||
}
|
||||
elif provider_type == "pytorch-cpu":
|
||||
return {
|
||||
"type": provider_type,
|
||||
"name": "PyTorch CPU",
|
||||
"installed": provider_type in await self.list_installed(),
|
||||
"size_mb": 300,
|
||||
}
|
||||
elif provider_type == "pytorch-cuda":
|
||||
return {
|
||||
"type": provider_type,
|
||||
"name": "PyTorch CUDA",
|
||||
"installed": provider_type in await self.list_installed(),
|
||||
"size_mb": 2400,
|
||||
}
|
||||
else:
|
||||
return {
|
||||
"type": provider_type,
|
||||
"name": provider_type,
|
||||
"installed": False,
|
||||
"size_mb": None,
|
||||
}
|
||||
|
||||
|
||||
def _get_free_port(self) -> int:
|
||||
"""Get a free port for the provider server."""
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
s.bind(('', 0))
|
||||
return s.getsockname()[1]
|
||||
|
||||
async def _wait_for_provider_health(self, base_url: str, timeout: int = 30) -> None:
|
||||
"""Wait for provider to become healthy."""
|
||||
import httpx
|
||||
import asyncio
|
||||
|
||||
start_time = asyncio.get_event_loop().time()
|
||||
while True:
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=2.0) as client:
|
||||
response = await client.get(f"{base_url}/tts/health")
|
||||
if response.status_code == 200:
|
||||
return
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if asyncio.get_event_loop().time() - start_time > timeout:
|
||||
raise TimeoutError(f"Provider did not become healthy within {timeout} seconds")
|
||||
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
|
||||
# Global provider manager instance
|
||||
_provider_manager: Optional[ProviderManager] = None
|
||||
|
||||
|
||||
def get_provider_manager() -> ProviderManager:
|
||||
"""Get the global provider manager instance."""
|
||||
global _provider_manager
|
||||
if _provider_manager is None:
|
||||
_provider_manager = ProviderManager()
|
||||
return _provider_manager
|
||||
@@ -0,0 +1,97 @@
|
||||
"""
|
||||
Base protocol for TTS providers.
|
||||
"""
|
||||
|
||||
from typing import Protocol, Optional, Tuple
|
||||
from typing_extensions import runtime_checkable
|
||||
import numpy as np
|
||||
|
||||
from .types import ProviderHealth, ProviderStatus
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class TTSProvider(Protocol):
|
||||
"""Protocol for TTS provider implementations."""
|
||||
|
||||
async def generate(
|
||||
self,
|
||||
text: str,
|
||||
voice_prompt: dict,
|
||||
language: str = "en",
|
||||
seed: Optional[int] = None,
|
||||
instruct: Optional[str] = None,
|
||||
) -> Tuple[np.ndarray, int]:
|
||||
"""
|
||||
Generate speech audio from text.
|
||||
|
||||
Args:
|
||||
text: Text to synthesize
|
||||
voice_prompt: Voice prompt dictionary
|
||||
language: Language code
|
||||
seed: Random seed for reproducibility
|
||||
instruct: Delivery instructions
|
||||
|
||||
Returns:
|
||||
Tuple of (audio_array, sample_rate)
|
||||
"""
|
||||
...
|
||||
|
||||
async def create_voice_prompt(
|
||||
self,
|
||||
audio_path: str,
|
||||
reference_text: str,
|
||||
use_cache: bool = True,
|
||||
) -> Tuple[dict, bool]:
|
||||
"""
|
||||
Create voice prompt from reference audio.
|
||||
|
||||
Args:
|
||||
audio_path: Path to reference audio file
|
||||
reference_text: Transcript of the audio
|
||||
use_cache: Whether to use cached prompts
|
||||
|
||||
Returns:
|
||||
Tuple of (voice_prompt_dict, was_cached)
|
||||
"""
|
||||
...
|
||||
|
||||
async def combine_voice_prompts(
|
||||
self,
|
||||
audio_paths: list[str],
|
||||
reference_texts: list[str],
|
||||
) -> Tuple[np.ndarray, str]:
|
||||
"""
|
||||
Combine multiple voice prompts.
|
||||
|
||||
Args:
|
||||
audio_paths: List of audio file paths
|
||||
reference_texts: List of reference texts
|
||||
|
||||
Returns:
|
||||
Tuple of (combined_audio_array, combined_text)
|
||||
"""
|
||||
...
|
||||
|
||||
async def load_model(self, model_size: str) -> None:
|
||||
"""Load TTS model."""
|
||||
...
|
||||
|
||||
def unload_model(self) -> None:
|
||||
"""Unload model to free memory."""
|
||||
...
|
||||
|
||||
def is_loaded(self) -> bool:
|
||||
"""Check if model is loaded."""
|
||||
...
|
||||
|
||||
def _get_model_path(self, model_size: str) -> str:
|
||||
"""Get model path for a given size."""
|
||||
...
|
||||
|
||||
async def health(self) -> ProviderHealth:
|
||||
"""Get provider health status."""
|
||||
...
|
||||
|
||||
async def status(self) -> ProviderStatus:
|
||||
"""Get provider model status."""
|
||||
...
|
||||
@@ -0,0 +1,139 @@
|
||||
"""
|
||||
Bundled provider that wraps existing MLX/PyTorch backends.
|
||||
"""
|
||||
|
||||
from typing import Optional, Tuple
|
||||
import numpy as np
|
||||
import platform
|
||||
|
||||
from .base import TTSProvider
|
||||
from .types import ProviderHealth, ProviderStatus
|
||||
from ..backends import get_tts_backend, TTSBackend
|
||||
from ..platform_detect import get_backend_type
|
||||
|
||||
|
||||
class BundledProvider:
|
||||
"""Provider that wraps the existing bundled TTS backend."""
|
||||
|
||||
def __init__(self):
|
||||
self._backend: Optional[TTSBackend] = None
|
||||
|
||||
def _get_backend(self) -> TTSBackend:
|
||||
"""Get or create backend instance."""
|
||||
if self._backend is None:
|
||||
self._backend = get_tts_backend()
|
||||
return self._backend
|
||||
|
||||
async def generate(
|
||||
self,
|
||||
text: str,
|
||||
voice_prompt: dict,
|
||||
language: str = "en",
|
||||
seed: Optional[int] = None,
|
||||
instruct: Optional[str] = None,
|
||||
) -> Tuple[np.ndarray, int]:
|
||||
"""Generate speech audio."""
|
||||
backend = self._get_backend()
|
||||
return await backend.generate(text, voice_prompt, language, seed, instruct)
|
||||
|
||||
async def create_voice_prompt(
|
||||
self,
|
||||
audio_path: str,
|
||||
reference_text: str,
|
||||
use_cache: bool = True,
|
||||
) -> Tuple[dict, bool]:
|
||||
"""Create voice prompt from reference audio."""
|
||||
backend = self._get_backend()
|
||||
return await backend.create_voice_prompt(audio_path, reference_text, use_cache)
|
||||
|
||||
async def combine_voice_prompts(
|
||||
self,
|
||||
audio_paths: list[str],
|
||||
reference_texts: list[str],
|
||||
) -> Tuple[np.ndarray, str]:
|
||||
"""Combine multiple voice prompts."""
|
||||
backend = self._get_backend()
|
||||
return await backend.combine_voice_prompts(audio_paths, reference_texts)
|
||||
|
||||
async def load_model(self, model_size: str) -> None:
|
||||
"""Load TTS model."""
|
||||
backend = self._get_backend()
|
||||
# Backends use load_model_async, but Protocol defines load_model
|
||||
if hasattr(backend, 'load_model_async'):
|
||||
await backend.load_model_async(model_size)
|
||||
else:
|
||||
await backend.load_model(model_size)
|
||||
|
||||
def unload_model(self) -> None:
|
||||
"""Unload model to free memory."""
|
||||
backend = self._get_backend()
|
||||
backend.unload_model()
|
||||
|
||||
def is_loaded(self) -> bool:
|
||||
"""Check if model is loaded."""
|
||||
backend = self._get_backend()
|
||||
return backend.is_loaded()
|
||||
|
||||
def _get_model_path(self, model_size: str) -> str:
|
||||
"""Get model path for a given size."""
|
||||
backend = self._get_backend()
|
||||
return backend._get_model_path(model_size)
|
||||
|
||||
async def health(self) -> ProviderHealth:
|
||||
"""Get provider health status."""
|
||||
backend = self._get_backend()
|
||||
backend_type = get_backend_type()
|
||||
|
||||
model_size = None
|
||||
if backend.is_loaded():
|
||||
# Try to get current model size from backend
|
||||
if hasattr(backend, '_current_model_size') and backend._current_model_size:
|
||||
model_size = backend._current_model_size
|
||||
|
||||
device = None
|
||||
if backend_type == "mlx":
|
||||
device = "metal"
|
||||
elif hasattr(backend, 'device'):
|
||||
device = backend.device
|
||||
|
||||
return ProviderHealth(
|
||||
status="healthy",
|
||||
provider=f"bundled-{backend_type}",
|
||||
version=None, # Provider versioning not implemented yet
|
||||
model=model_size,
|
||||
device=device,
|
||||
)
|
||||
|
||||
async def status(self) -> ProviderStatus:
|
||||
"""Get provider model status."""
|
||||
backend = self._get_backend()
|
||||
backend_type = get_backend_type()
|
||||
|
||||
model_size = None
|
||||
if backend.is_loaded():
|
||||
if hasattr(backend, '_current_model_size') and backend._current_model_size:
|
||||
model_size = backend._current_model_size
|
||||
|
||||
available_sizes = ["1.7B"]
|
||||
if backend_type == "pytorch":
|
||||
available_sizes.append("0.6B")
|
||||
|
||||
gpu_available = None
|
||||
vram_used_mb = None
|
||||
|
||||
if backend_type == "pytorch":
|
||||
try:
|
||||
import torch
|
||||
gpu_available = torch.cuda.is_available()
|
||||
if gpu_available:
|
||||
vram_used_mb = torch.cuda.memory_allocated() / 1024 / 1024
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
return ProviderStatus(
|
||||
model_loaded=backend.is_loaded(),
|
||||
model_size=model_size,
|
||||
available_sizes=available_sizes,
|
||||
gpu_available=gpu_available,
|
||||
vram_used_mb=int(vram_used_mb) if vram_used_mb else None,
|
||||
)
|
||||
@@ -0,0 +1,211 @@
|
||||
"""
|
||||
Provider download and installation manager.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import httpx
|
||||
import platform
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
from .types import ProviderType
|
||||
from ..utils.progress import get_progress_manager
|
||||
from ..utils.tasks import get_task_manager
|
||||
|
||||
|
||||
# Provider version (independent of app version)
|
||||
PROVIDER_VERSION = "1.0.0"
|
||||
|
||||
# Base URL for provider downloads (Cloudflare R2)
|
||||
PROVIDER_DOWNLOAD_BASE_URL = "https://downloads.voicebox.sh/providers"
|
||||
|
||||
|
||||
def _get_providers_dir() -> Path:
|
||||
"""Get the directory where providers are stored."""
|
||||
system = platform.system()
|
||||
|
||||
if system == "Windows":
|
||||
appdata = Path.home() / "AppData" / "Roaming"
|
||||
elif system == "Darwin":
|
||||
appdata = Path.home() / "Library" / "Application Support"
|
||||
else: # Linux
|
||||
appdata = Path.home() / ".local" / "share"
|
||||
|
||||
providers_dir = appdata / "voicebox" / "providers"
|
||||
providers_dir.mkdir(parents=True, exist_ok=True)
|
||||
return providers_dir
|
||||
|
||||
|
||||
def _get_provider_binary_name(provider_type: str) -> str:
|
||||
"""Get the local binary filename for a provider type."""
|
||||
system = platform.system()
|
||||
ext = ".exe" if system == "Windows" else ""
|
||||
|
||||
binary_map = {
|
||||
"pytorch-cpu": f"tts-provider-pytorch-cpu{ext}",
|
||||
"pytorch-cuda": f"tts-provider-pytorch-cuda{ext}",
|
||||
}
|
||||
|
||||
if provider_type not in binary_map:
|
||||
raise ValueError(f"Unknown provider type: {provider_type}")
|
||||
|
||||
return binary_map[provider_type]
|
||||
|
||||
|
||||
def _get_provider_download_name(provider_type: str) -> str:
|
||||
"""Get the remote download filename for a provider type (includes platform suffix)."""
|
||||
system = platform.system()
|
||||
|
||||
if system == "Windows":
|
||||
platform_suffix = "windows"
|
||||
ext = ".exe"
|
||||
elif system == "Linux":
|
||||
platform_suffix = "linux"
|
||||
ext = ""
|
||||
else:
|
||||
raise ValueError(f"Provider downloads not supported on {system}")
|
||||
|
||||
return f"tts-provider-{provider_type}-{platform_suffix}{ext}"
|
||||
|
||||
|
||||
def _get_provider_download_url(provider_type: str) -> str:
|
||||
"""Get the download URL for a provider."""
|
||||
download_name = _get_provider_download_name(provider_type)
|
||||
return f"{PROVIDER_DOWNLOAD_BASE_URL}/v{PROVIDER_VERSION}/{download_name}"
|
||||
|
||||
|
||||
async def download_provider(provider_type: str) -> Path:
|
||||
"""
|
||||
Download a provider binary from Cloudflare R2.
|
||||
|
||||
Args:
|
||||
provider_type: Type of provider to download (e.g., "pytorch-cpu")
|
||||
|
||||
Returns:
|
||||
Path to the downloaded provider binary
|
||||
|
||||
Raises:
|
||||
ValueError: If provider_type is invalid
|
||||
httpx.HTTPError: If download fails
|
||||
"""
|
||||
if provider_type not in ["pytorch-cpu", "pytorch-cuda"]:
|
||||
raise ValueError(f"Provider type {provider_type} cannot be downloaded")
|
||||
|
||||
progress_manager = get_progress_manager()
|
||||
task_manager = get_task_manager()
|
||||
|
||||
binary_name = _get_provider_binary_name(provider_type)
|
||||
download_url = _get_provider_download_url(provider_type)
|
||||
destination = _get_providers_dir() / binary_name
|
||||
|
||||
# Start tracking download
|
||||
task_manager.start_download(provider_type)
|
||||
|
||||
# Initialize progress state
|
||||
progress_manager.update_progress(
|
||||
model_name=provider_type,
|
||||
current=0,
|
||||
total=0, # Will be updated once we get Content-Length
|
||||
filename=binary_name,
|
||||
status="downloading",
|
||||
)
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=300.0) as client:
|
||||
# First, get the file size
|
||||
async with client.stream("GET", download_url) as response:
|
||||
response.raise_for_status()
|
||||
|
||||
# Get total size from Content-Length header
|
||||
total_size = int(response.headers.get("Content-Length", 0))
|
||||
|
||||
if total_size > 0:
|
||||
progress_manager.update_progress(
|
||||
model_name=provider_type,
|
||||
current=0,
|
||||
total=total_size,
|
||||
filename=binary_name,
|
||||
status="downloading",
|
||||
)
|
||||
|
||||
# Download with progress tracking
|
||||
downloaded = 0
|
||||
with open(destination, "wb") as f:
|
||||
async for chunk in response.aiter_bytes(chunk_size=8192):
|
||||
f.write(chunk)
|
||||
downloaded += len(chunk)
|
||||
|
||||
# Update progress
|
||||
progress_manager.update_progress(
|
||||
model_name=provider_type,
|
||||
current=downloaded,
|
||||
total=total_size if total_size > 0 else downloaded,
|
||||
filename=binary_name,
|
||||
status="downloading",
|
||||
)
|
||||
|
||||
# Mark as complete
|
||||
progress_manager.update_progress(
|
||||
model_name=provider_type,
|
||||
current=downloaded,
|
||||
total=downloaded,
|
||||
filename=binary_name,
|
||||
status="complete",
|
||||
)
|
||||
task_manager.complete_download(provider_type)
|
||||
|
||||
# Make executable on Unix systems
|
||||
if platform.system() != "Windows":
|
||||
destination.chmod(0o755)
|
||||
|
||||
return destination
|
||||
|
||||
except Exception as e:
|
||||
# Mark as error
|
||||
progress_manager.update_progress(
|
||||
model_name=provider_type,
|
||||
current=0,
|
||||
total=0,
|
||||
filename=binary_name,
|
||||
status="error",
|
||||
)
|
||||
task_manager.error_download(provider_type, str(e))
|
||||
raise
|
||||
|
||||
|
||||
def get_provider_binary_path(provider_type: str) -> Optional[Path]:
|
||||
"""
|
||||
Get the path to an installed provider binary.
|
||||
|
||||
Args:
|
||||
provider_type: Type of provider
|
||||
|
||||
Returns:
|
||||
Path to provider binary, or None if not installed
|
||||
"""
|
||||
binary_name = _get_provider_binary_name(provider_type)
|
||||
provider_path = _get_providers_dir() / binary_name
|
||||
|
||||
if provider_path.exists() and provider_path.is_file():
|
||||
return provider_path
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def delete_provider(provider_type: str) -> bool:
|
||||
"""
|
||||
Delete an installed provider binary.
|
||||
|
||||
Args:
|
||||
provider_type: Type of provider to delete
|
||||
|
||||
Returns:
|
||||
True if deleted, False if not found
|
||||
"""
|
||||
provider_path = get_provider_binary_path(provider_type)
|
||||
|
||||
if provider_path and provider_path.exists():
|
||||
provider_path.unlink()
|
||||
return True
|
||||
|
||||
return False
|
||||
@@ -0,0 +1,187 @@
|
||||
"""
|
||||
Local provider that communicates with standalone provider servers via HTTP.
|
||||
"""
|
||||
|
||||
from typing import Optional, Tuple
|
||||
import base64
|
||||
import io
|
||||
import numpy as np
|
||||
import httpx
|
||||
import soundfile as sf
|
||||
|
||||
from .base import TTSProvider
|
||||
from .types import ProviderHealth, ProviderStatus
|
||||
|
||||
|
||||
class LocalProvider:
|
||||
"""Provider that communicates with local subprocess via HTTP."""
|
||||
|
||||
def __init__(self, base_url: str):
|
||||
"""
|
||||
Initialize local provider.
|
||||
|
||||
Args:
|
||||
base_url: Base URL of the provider server (e.g., "http://localhost:8000")
|
||||
"""
|
||||
self.base_url = base_url.rstrip('/')
|
||||
self.client = httpx.AsyncClient(timeout=300.0) # 5 minute timeout for generation
|
||||
|
||||
async def generate(
|
||||
self,
|
||||
text: str,
|
||||
voice_prompt: dict,
|
||||
language: str = "en",
|
||||
seed: Optional[int] = None,
|
||||
instruct: Optional[str] = None,
|
||||
) -> Tuple[np.ndarray, int]:
|
||||
"""Generate speech audio."""
|
||||
response = await self.client.post(
|
||||
f"{self.base_url}/tts/generate",
|
||||
json={
|
||||
"text": text,
|
||||
"voice_prompt": voice_prompt,
|
||||
"language": language,
|
||||
"seed": seed,
|
||||
"model_size": "1.7B", # TODO: Make configurable
|
||||
}
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
# Decode base64 audio
|
||||
audio_bytes = base64.b64decode(data["audio"])
|
||||
audio_buffer = io.BytesIO(audio_bytes)
|
||||
audio, sample_rate = sf.read(audio_buffer)
|
||||
|
||||
return audio, data["sample_rate"]
|
||||
|
||||
async def create_voice_prompt(
|
||||
self,
|
||||
audio_path: str,
|
||||
reference_text: str,
|
||||
use_cache: bool = True,
|
||||
) -> Tuple[dict, bool]:
|
||||
"""Create voice prompt from reference audio."""
|
||||
# Read audio file
|
||||
with open(audio_path, 'rb') as f:
|
||||
audio_data = f.read()
|
||||
|
||||
# Send multipart form data
|
||||
files = {
|
||||
"audio": ("audio.wav", audio_data, "audio/wav")
|
||||
}
|
||||
data = {
|
||||
"reference_text": reference_text,
|
||||
"use_cache": str(use_cache).lower(),
|
||||
}
|
||||
|
||||
response = await self.client.post(
|
||||
f"{self.base_url}/tts/create_voice_prompt",
|
||||
files=files,
|
||||
data=data,
|
||||
)
|
||||
response.raise_for_status()
|
||||
result = response.json()
|
||||
|
||||
return result["voice_prompt"], result.get("was_cached", False)
|
||||
|
||||
async def combine_voice_prompts(
|
||||
self,
|
||||
audio_paths: list[str],
|
||||
reference_texts: list[str],
|
||||
) -> Tuple[np.ndarray, str]:
|
||||
"""
|
||||
Combine multiple voice prompts.
|
||||
|
||||
Note: This is not implemented in the provider API yet.
|
||||
For now, we'll combine locally by concatenating audio.
|
||||
"""
|
||||
import numpy as np
|
||||
from ..utils.audio import load_audio, normalize_audio
|
||||
|
||||
combined_audio = []
|
||||
for audio_path in audio_paths:
|
||||
audio, sr = load_audio(audio_path)
|
||||
audio = normalize_audio(audio)
|
||||
combined_audio.append(audio)
|
||||
|
||||
# Concatenate audio
|
||||
mixed = np.concatenate(combined_audio)
|
||||
mixed = normalize_audio(mixed)
|
||||
|
||||
# Combine texts
|
||||
combined_text = " ".join(reference_texts)
|
||||
|
||||
return mixed, combined_text
|
||||
|
||||
async def load_model(self, model_size: str) -> None:
|
||||
"""Load TTS model."""
|
||||
# Model loading is handled automatically by the provider server
|
||||
# when generate() is called, so this is a no-op
|
||||
pass
|
||||
|
||||
def unload_model(self) -> None:
|
||||
"""Unload model to free memory."""
|
||||
# Model unloading is handled by the provider server
|
||||
# This is a no-op for local providers
|
||||
pass
|
||||
|
||||
def is_loaded(self) -> bool:
|
||||
"""Check if model is loaded."""
|
||||
# We can't know this without querying the provider
|
||||
# Return True optimistically
|
||||
return True
|
||||
|
||||
def _get_model_path(self, model_size: str) -> str:
|
||||
"""Get model path for a given size."""
|
||||
# For local providers, model paths are handled by the provider server
|
||||
# Return a placeholder
|
||||
return f"Qwen/Qwen3-TTS-12Hz-{model_size}-Base"
|
||||
|
||||
async def health(self) -> ProviderHealth:
|
||||
"""Get provider health status."""
|
||||
try:
|
||||
response = await self.client.get(f"{self.base_url}/tts/health")
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
return ProviderHealth(
|
||||
status=data["status"],
|
||||
provider=data["provider"],
|
||||
version=data.get("version"),
|
||||
model=data.get("model"),
|
||||
device=data.get("device"),
|
||||
)
|
||||
except Exception as e:
|
||||
return ProviderHealth(
|
||||
status="unhealthy",
|
||||
provider="local",
|
||||
version=None,
|
||||
model=None,
|
||||
device=None,
|
||||
)
|
||||
|
||||
async def status(self) -> ProviderStatus:
|
||||
"""Get provider model status."""
|
||||
try:
|
||||
response = await self.client.get(f"{self.base_url}/tts/status")
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
return ProviderStatus(
|
||||
model_loaded=data["model_loaded"],
|
||||
model_size=data.get("model_size"),
|
||||
available_sizes=data.get("available_sizes", []),
|
||||
gpu_available=data.get("gpu_available"),
|
||||
vram_used_mb=data.get("vram_used_mb"),
|
||||
)
|
||||
except Exception as e:
|
||||
return ProviderStatus(
|
||||
model_loaded=False,
|
||||
model_size=None,
|
||||
available_sizes=[],
|
||||
gpu_available=None,
|
||||
vram_used_mb=None,
|
||||
)
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""Stop the provider (close HTTP client)."""
|
||||
await self.client.aclose()
|
||||
@@ -0,0 +1,34 @@
|
||||
"""
|
||||
Shared types for TTS providers.
|
||||
"""
|
||||
|
||||
from typing import Optional, TypedDict
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class ProviderType(str, Enum):
|
||||
"""Available provider types."""
|
||||
BUNDLED_MLX = "bundled-mlx"
|
||||
BUNDLED_PYTORCH = "bundled-pytorch"
|
||||
PYTORCH_CPU = "pytorch-cpu"
|
||||
PYTORCH_CUDA = "pytorch-cuda"
|
||||
REMOTE = "remote"
|
||||
OPENAI = "openai"
|
||||
|
||||
|
||||
class ProviderHealth(TypedDict):
|
||||
"""Provider health status."""
|
||||
status: str # "healthy", "unhealthy", "starting"
|
||||
provider: str
|
||||
version: Optional[str]
|
||||
model: Optional[str]
|
||||
device: Optional[str]
|
||||
|
||||
|
||||
class ProviderStatus(TypedDict):
|
||||
"""Provider model status."""
|
||||
model_loaded: bool
|
||||
model_size: Optional[str]
|
||||
available_sizes: list[str]
|
||||
gpu_available: Optional[bool]
|
||||
vram_used_mb: Optional[int]
|
||||
+36
-16
@@ -1,5 +1,5 @@
|
||||
"""
|
||||
TTS inference module - delegates to backend abstraction layer.
|
||||
TTS inference module - delegates to provider abstraction layer.
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
@@ -7,31 +7,51 @@ import numpy as np
|
||||
import io
|
||||
import soundfile as sf
|
||||
|
||||
from .backends import get_tts_backend, TTSBackend
|
||||
from .backends import TTSBackend
|
||||
from .providers import get_provider_manager
|
||||
from .providers.base import TTSProvider
|
||||
|
||||
|
||||
def get_tts_model() -> TTSBackend:
|
||||
def get_tts_model() -> TTSProvider:
|
||||
"""
|
||||
Get TTS backend instance (MLX or PyTorch based on platform).
|
||||
Get TTS provider instance (via ProviderManager).
|
||||
|
||||
Returns:
|
||||
TTS backend instance
|
||||
TTS provider instance
|
||||
"""
|
||||
return get_tts_backend()
|
||||
manager = get_provider_manager()
|
||||
# Note: This is async but we need sync interface for backward compatibility
|
||||
# In practice, this will be called from async contexts
|
||||
import asyncio
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
if loop.is_running():
|
||||
# We're in an async context, but can't await here
|
||||
# Return a wrapper that will use the provider manager
|
||||
return manager._get_default_provider()
|
||||
else:
|
||||
return loop.run_until_complete(manager.get_active_provider())
|
||||
except RuntimeError:
|
||||
# No event loop, return default
|
||||
return manager._get_default_provider()
|
||||
|
||||
|
||||
async def get_tts_model_async() -> TTSProvider:
|
||||
"""
|
||||
Get TTS provider instance asynchronously.
|
||||
|
||||
Returns:
|
||||
TTS provider instance
|
||||
"""
|
||||
manager = get_provider_manager()
|
||||
return await manager.get_active_provider()
|
||||
|
||||
|
||||
def unload_tts_model():
|
||||
"""Unload TTS model to free memory."""
|
||||
backend = get_tts_backend()
|
||||
backend.unload_model()
|
||||
|
||||
|
||||
def audio_to_wav_bytes(audio: np.ndarray, sample_rate: int) -> bytes:
|
||||
"""Convert audio array to WAV bytes."""
|
||||
buffer = io.BytesIO()
|
||||
sf.write(buffer, audio, sample_rate, format="WAV")
|
||||
buffer.seek(0)
|
||||
return buffer.read()
|
||||
manager = get_provider_manager()
|
||||
provider = manager._get_default_provider()
|
||||
provider.unload_model()
|
||||
|
||||
|
||||
def audio_to_wav_bytes(audio: np.ndarray, sample_rate: int) -> bytes:
|
||||
|
||||
Reference in New Issue
Block a user