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:
Jamie Pine
2026-01-31 03:05:50 -08:00
parent 220333b3bb
commit 80689ad8ce
21 changed files with 2811 additions and 171 deletions
+15 -17
View File
@@ -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
View File
@@ -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"]:
+220
View File
@@ -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
+97
View File
@@ -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."""
...
+139
View File
@@ -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,
)
+211
View File
@@ -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
+187
View File
@@ -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()
+34
View File
@@ -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
View File
@@ -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: