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
+132 -12
View File
@@ -6,7 +6,114 @@ on:
tags:
- "v*"
env:
PROVIDER_VERSION: "1.0.0"
jobs:
# ============================================
# Build TTS Providers (uploaded to R2, not GitHub)
# ============================================
build-providers:
runs-on: ${{ matrix.platform }}
strategy:
fail-fast: false
matrix:
include:
# PyTorch CPU provider (Windows)
- platform: "windows-latest"
provider: "pytorch-cpu"
python-version: "3.12"
# PyTorch CUDA provider (Windows) - large binary, uploaded to R2
- platform: "windows-latest"
provider: "pytorch-cuda"
python-version: "3.12"
# PyTorch CPU provider (Linux)
- platform: "ubuntu-22.04"
provider: "pytorch-cpu"
python-version: "3.12"
# PyTorch CUDA provider (Linux) - large binary, uploaded to R2
- platform: "ubuntu-22.04"
provider: "pytorch-cuda"
python-version: "3.12"
steps:
- uses: actions/checkout@v4
- name: Install dependencies (ubuntu only)
if: matrix.platform == 'ubuntu-22.04'
run: |
sudo apt-get update
sudo apt-get install -y llvm-dev
- name: Setup Python
uses: actions/setup-python@v5
with:
python-version: ${{ matrix.python-version }}
cache: "pip"
- name: Install Python dependencies (CPU)
if: matrix.provider == 'pytorch-cpu'
run: |
python -m pip install --upgrade pip
pip install pyinstaller
pip install -r providers/pytorch-cpu/requirements.txt
pip install -r backend/requirements.txt
- name: Install Python dependencies (CUDA)
if: matrix.provider == 'pytorch-cuda'
run: |
python -m pip install --upgrade pip
pip install pyinstaller
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121
pip install -r providers/pytorch-cuda/requirements.txt
pip install -r backend/requirements.txt
- name: Build provider binary
shell: bash
run: |
cd providers/${{ matrix.provider }}
python build.py
- name: Upload provider to R2
shell: bash
env:
R2_ACCESS_KEY_ID: ${{ secrets.R2_ACCESS_KEY_ID }}
R2_SECRET_ACCESS_KEY: ${{ secrets.R2_SECRET_ACCESS_KEY }}
R2_ENDPOINT: ${{ secrets.R2_ENDPOINT }}
run: |
# Install AWS CLI (compatible with R2)
pip install awscli
# Configure AWS CLI for R2
aws configure set aws_access_key_id $R2_ACCESS_KEY_ID
aws configure set aws_secret_access_key $R2_SECRET_ACCESS_KEY
aws configure set region auto
# Determine binary name based on platform
if [ "${{ matrix.platform }}" == "windows-latest" ]; then
BINARY_NAME="tts-provider-${{ matrix.provider }}.exe"
BINARY_PATH="providers/${{ matrix.provider }}/dist/tts-provider-${{ matrix.provider }}.exe"
else
BINARY_NAME="tts-provider-${{ matrix.provider }}"
BINARY_PATH="providers/${{ matrix.provider }}/dist/tts-provider-${{ matrix.provider }}"
fi
# Add platform suffix for clarity
if [ "${{ matrix.platform }}" == "windows-latest" ]; then
UPLOAD_NAME="tts-provider-${{ matrix.provider }}-windows.exe"
else
UPLOAD_NAME="tts-provider-${{ matrix.provider }}-linux"
fi
# Upload to R2 (bucket: voicebox)
aws s3 cp "$BINARY_PATH" "s3://voicebox/providers/v${{ env.PROVIDER_VERSION }}/$UPLOAD_NAME" \
--endpoint-url "$R2_ENDPOINT"
echo "Uploaded $UPLOAD_NAME to R2"
# ============================================
# Build Main App (without bundled TTS on Win/Linux)
# ============================================
release:
permissions:
contents: write
@@ -14,22 +121,26 @@ jobs:
fail-fast: false
matrix:
include:
# macOS Apple Silicon - MLX bundled (works out of the box)
- platform: "macos-latest"
args: "--target aarch64-apple-darwin"
python-version: "3.12"
backend: "mlx"
# macOS Intel - PyTorch bundled (smaller user base, keep simple)
- platform: "macos-15-intel"
args: "--target x86_64-apple-darwin"
python-version: "3.12"
backend: "pytorch"
# Linux - No TTS bundled, providers downloaded separately
# - platform: 'ubuntu-22.04'
# args: ''
# python-version: '3.12'
# backend: 'pytorch'
# backend: 'none'
# Windows - No TTS bundled, providers downloaded separately
- platform: "windows-latest"
args: ""
python-version: "3.12"
backend: "pytorch"
backend: "none"
runs-on: ${{ matrix.platform }}
@@ -55,23 +166,27 @@ jobs:
python-version: ${{ matrix.python-version }}
cache: "pip"
- name: Install Python dependencies
- name: Install Python dependencies (with TTS)
if: matrix.backend != 'none'
run: |
python -m pip install --upgrade pip
pip install pyinstaller
pip install -r backend/requirements.txt
- name: Install Python dependencies (without TTS)
if: matrix.backend == 'none'
run: |
python -m pip install --upgrade pip
pip install pyinstaller
# Install base requirements without PyTorch/Qwen-TTS
pip install fastapi uvicorn sqlalchemy librosa soundfile numpy httpx
pip install huggingface_hub # For Whisper downloads
- name: Install MLX dependencies (Apple Silicon only)
if: matrix.backend == 'mlx'
run: |
pip install -r backend/requirements-mlx.txt
# - name: Install PyTorch with CUDA (Windows only)
# if: matrix.platform == 'windows-latest'
# run: |
# pip install torch --index-url https://download.pytorch.org/whl/cu121 --force-reinstall --no-deps
# pip install torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121
- name: Build Python server (Linux/macOS)
if: matrix.platform != 'windows-latest'
run: |
@@ -148,10 +263,15 @@ jobs:
See the assets below to download and install this version.
### Installation
- **macOS (Apple Silicon)**: Download the `aarch64.dmg` file - uses MLX for fast native inference
- **macOS (Apple Silicon)**: Download the `aarch64.dmg` file - uses MLX for fast native inference (works out of the box)
- **macOS (Intel)**: Download the `x64.dmg` file - uses PyTorch
- **Windows**: Download the `.msi` installer
- **Linux**: Download the `.AppImage` or `.deb` package
- **Windows**: Download the `.msi` installer - requires downloading a TTS provider on first use
- **Linux**: Download the `.AppImage` or `.deb` package - requires downloading a TTS provider on first use
### TTS Providers (Windows/Linux)
Windows and Linux users will be prompted to download a TTS provider on first launch:
- **PyTorch CPU** (~300MB) - Works on any system
- **PyTorch CUDA** (~2.4GB) - 4-5x faster on NVIDIA GPUs
The app includes automatic updates - future updates will be installed automatically.
releaseDraft: true
@@ -0,0 +1,395 @@
import { useMutation, useQuery, useQueryClient } from '@tanstack/react-query';
import { Download, Loader2, Trash2 } from 'lucide-react';
import { useCallback, useState } from 'react';
import {
AlertDialog,
AlertDialogAction,
AlertDialogCancel,
AlertDialogContent,
AlertDialogDescription,
AlertDialogFooter,
AlertDialogHeader,
AlertDialogTitle,
} from '@/components/ui/alert-dialog';
import { Badge } from '@/components/ui/badge';
import { Button } from '@/components/ui/button';
import { Card, CardContent, CardDescription, CardHeader, CardTitle } from '@/components/ui/card';
import { Label } from '@/components/ui/label';
import { RadioGroup, RadioGroupItem } from '@/components/ui/radio-group';
import { useToast } from '@/components/ui/use-toast';
import { apiClient } from '@/lib/api/client';
import { useModelDownloadToast } from '@/lib/hooks/useModelDownloadToast';
const isMacOS = () => navigator.platform.toLowerCase().includes('mac');
type ProviderType = 'auto' | 'bundled-mlx' | 'bundled-pytorch' | 'pytorch-cpu' | 'pytorch-cuda' | 'remote' | 'openai';
export function ProviderSettings() {
const { toast } = useToast();
const queryClient = useQueryClient();
const [selectedProvider, setSelectedProvider] = useState<ProviderType>('auto');
const [downloadingProvider, setDownloadingProvider] = useState<string | null>(null);
const { data: providersData, isLoading } = useQuery({
queryKey: ['providers'],
queryFn: async () => {
return await apiClient.listProviders();
},
refetchInterval: 5000,
});
const { data: activeProvider } = useQuery({
queryKey: ['activeProvider'],
queryFn: async () => {
return await apiClient.getActiveProvider();
},
refetchInterval: 5000,
});
// Callbacks for download completion
const handleDownloadComplete = useCallback(() => {
setDownloadingProvider(null);
queryClient.invalidateQueries({ queryKey: ['providers'] });
}, [queryClient]);
const handleDownloadError = useCallback(() => {
setDownloadingProvider(null);
}, []);
// Use progress toast hook for the downloading provider
useModelDownloadToast({
modelName: downloadingProvider || '',
displayName: downloadingProvider || '',
enabled: !!downloadingProvider,
onComplete: handleDownloadComplete,
onError: handleDownloadError,
});
const [deleteDialogOpen, setDeleteDialogOpen] = useState(false);
const [providerToDelete, setProviderToDelete] = useState<string | null>(null);
const downloadMutation = useMutation({
mutationFn: async (providerType: string) => {
return await apiClient.downloadProvider(providerType);
},
onSuccess: (_, providerType) => {
setDownloadingProvider(providerType);
queryClient.invalidateQueries({ queryKey: ['providers'] });
},
onError: (error: Error) => {
toast({
title: 'Download failed',
description: error.message,
variant: 'destructive',
});
},
});
const startMutation = useMutation({
mutationFn: async (providerType: string) => {
return await apiClient.startProvider(providerType);
},
onSuccess: () => {
queryClient.invalidateQueries({ queryKey: ['activeProvider'] });
toast({
title: 'Provider started',
description: 'The provider has been started successfully',
});
},
onError: (error: Error) => {
toast({
title: 'Failed to start provider',
description: error.message,
variant: 'destructive',
});
},
});
const deleteMutation = useMutation({
mutationFn: async (providerType: string) => {
return await apiClient.deleteProvider(providerType);
},
onSuccess: () => {
queryClient.invalidateQueries({ queryKey: ['providers'] });
toast({
title: 'Provider deleted',
description: 'The provider has been deleted successfully',
});
},
onError: (error: Error) => {
toast({
title: 'Failed to delete provider',
description: error.message,
variant: 'destructive',
});
},
});
const handleDownload = async (providerType: string) => {
downloadMutation.mutate(providerType);
};
const handleStart = async (providerType: string) => {
startMutation.mutate(providerType);
};
const handleDelete = (providerType: string) => {
setProviderToDelete(providerType);
setDeleteDialogOpen(true);
};
const confirmDelete = () => {
if (providerToDelete) {
deleteMutation.mutate(providerToDelete);
setDeleteDialogOpen(false);
setProviderToDelete(null);
}
};
if (isLoading) {
return (
<Card>
<CardHeader>
<CardTitle>TTS Provider</CardTitle>
<CardDescription>Choose how Voicebox generates speech</CardDescription>
</CardHeader>
<CardContent>
<div className="flex items-center justify-center py-8">
<Loader2 className="h-6 w-6 animate-spin" />
</div>
</CardContent>
</Card>
);
}
const installedProviders = providersData?.installed || [];
// Determine current active provider
const currentProvider = activeProvider?.provider || 'auto';
return (
<>
<Card>
<CardHeader>
<CardTitle>TTS Provider</CardTitle>
<CardDescription>Choose how Voicebox generates speech</CardDescription>
</CardHeader>
<CardContent>
<RadioGroup
value={selectedProvider}
onValueChange={(value) => setSelectedProvider(value as ProviderType)}
>
{/* Auto-detect */}
<div className="flex items-center space-x-2 py-2">
<RadioGroupItem value="auto" id="auto" />
<Label htmlFor="auto" className="flex-1 cursor-pointer">
<div className="font-medium">Auto-detect (Recommended)</div>
<div className="text-sm text-muted-foreground">
Automatically choose the best available provider
</div>
</Label>
{currentProvider === 'auto' && (
<Badge variant="outline" className="ml-2">
Active
</Badge>
)}
</div>
{/* PyTorch CUDA */}
<div className="flex items-center justify-between py-2">
<div className="flex items-center space-x-2 flex-1">
<RadioGroupItem value="pytorch-cuda" id="cuda" />
<Label htmlFor="cuda" className="flex-1 cursor-pointer">
<div className="font-medium">PyTorch CUDA (NVIDIA GPU)</div>
<div className="text-sm text-muted-foreground">
4-5x faster inference on NVIDIA GPUs
</div>
</Label>
</div>
<div className="flex items-center gap-2">
{currentProvider === 'pytorch-cuda' && (
<Badge variant="outline">Active</Badge>
)}
{!installedProviders.includes('pytorch-cuda') && (
<Button
onClick={() => handleDownload('pytorch-cuda')}
size="sm"
disabled={downloadingProvider === 'pytorch-cuda'}
>
{downloadingProvider === 'pytorch-cuda' ? (
<Loader2 className="h-4 w-4 animate-spin" />
) : (
<>
<Download className="h-4 w-4 mr-1" />
Download (2.4GB)
</>
)}
</Button>
)}
{installedProviders.includes('pytorch-cuda') && selectedProvider !== 'pytorch-cuda' && (
<Button
onClick={() => handleStart('pytorch-cuda')}
size="sm"
variant="outline"
>
Start
</Button>
)}
{installedProviders.includes('pytorch-cuda') && (
<Button
onClick={() => handleDelete('pytorch-cuda')}
size="sm"
variant="ghost"
>
<Trash2 className="h-4 w-4" />
</Button>
)}
</div>
</div>
{/* PyTorch CPU (Windows/Linux only) */}
{!isMacOS() && (
<div className="flex items-center justify-between py-2">
<div className="flex items-center space-x-2 flex-1">
<RadioGroupItem value="pytorch-cpu" id="cpu" />
<Label htmlFor="cpu" className="flex-1 cursor-pointer">
<div className="font-medium">PyTorch CPU</div>
<div className="text-sm text-muted-foreground">
Works on any system, slower inference
</div>
</Label>
</div>
<div className="flex items-center gap-2">
{currentProvider === 'pytorch-cpu' && (
<Badge variant="outline">Active</Badge>
)}
{!installedProviders.includes('pytorch-cpu') && (
<Button
onClick={() => handleDownload('pytorch-cpu')}
size="sm"
disabled={downloadingProvider === 'pytorch-cpu'}
>
{downloadingProvider === 'pytorch-cpu' ? (
<Loader2 className="h-4 w-4 animate-spin" />
) : (
<>
<Download className="h-4 w-4 mr-1" />
Download (300MB)
</>
)}
</Button>
)}
{installedProviders.includes('pytorch-cpu') && selectedProvider !== 'pytorch-cpu' && (
<Button
onClick={() => handleStart('pytorch-cpu')}
size="sm"
variant="outline"
>
Start
</Button>
)}
{installedProviders.includes('pytorch-cpu') && (
<Button
onClick={() => handleDelete('pytorch-cpu')}
size="sm"
variant="ghost"
>
<Trash2 className="h-4 w-4" />
</Button>
)}
</div>
</div>
)}
{/* MLX bundled (macOS only) */}
{isMacOS() && (
<div className="p-3 bg-muted rounded-md">
<div className="text-sm">
<div className="font-medium flex items-center gap-2">
MLX (Apple Silicon)
{currentProvider === 'bundled-mlx' && (
<Badge variant="outline">Active</Badge>
)}
</div>
<div className="text-muted-foreground mt-1">
Bundled with the app - optimized for M1/M2/M3 chips
</div>
</div>
</div>
)}
{/* Remote */}
<div className="space-y-2 py-2">
<div className="flex items-center space-x-2">
<RadioGroupItem value="remote" id="remote" />
<Label htmlFor="remote" className="flex-1 cursor-pointer">
<div className="font-medium">Remote Server</div>
<div className="text-sm text-muted-foreground">
Connect to your own TTS server
</div>
</Label>
</div>
{selectedProvider === 'remote' && (
<div className="ml-6">
<input
type="text"
placeholder="http://your-server:8000"
className="w-full px-3 py-2 border rounded-md"
disabled
/>
<div className="text-xs text-muted-foreground mt-1">
Remote provider support coming soon
</div>
</div>
)}
</div>
{/* OpenAI */}
<div className="space-y-2 py-2">
<div className="flex items-center space-x-2">
<RadioGroupItem value="openai" id="openai" />
<Label htmlFor="openai" className="flex-1 cursor-pointer">
<div className="font-medium">OpenAI API</div>
<div className="text-sm text-muted-foreground">
Use OpenAI's TTS API (requires API key)
</div>
</Label>
</div>
{selectedProvider === 'openai' && (
<div className="ml-6">
<input
type="password"
placeholder="sk-..."
className="w-full px-3 py-2 border rounded-md"
disabled
/>
<div className="text-xs text-muted-foreground mt-1">
OpenAI provider support coming soon
</div>
</div>
)}
</div>
</RadioGroup>
</CardContent>
</Card>
<AlertDialog open={deleteDialogOpen} onOpenChange={setDeleteDialogOpen}>
<AlertDialogContent>
<AlertDialogHeader>
<AlertDialogTitle>Delete Provider</AlertDialogTitle>
<AlertDialogDescription>
Are you sure you want to delete {providerToDelete}? This will remove the provider
binary from your system. You can download it again later if needed.
</AlertDialogDescription>
</AlertDialogHeader>
<AlertDialogFooter>
<AlertDialogCancel>Cancel</AlertDialogCancel>
<AlertDialogAction onClick={confirmDelete} className="bg-destructive text-destructive-foreground">
Delete
</AlertDialogAction>
</AlertDialogFooter>
</AlertDialogContent>
</AlertDialog>
</>
);
}
@@ -1,6 +1,7 @@
import { ConnectionForm } from '@/components/ServerSettings/ConnectionForm';
import { ServerStatus } from '@/components/ServerSettings/ServerStatus';
import { UpdateStatus } from '@/components/ServerSettings/UpdateStatus';
import { ProviderSettings } from '@/components/ServerSettings/ProviderSettings';
import { usePlatform } from '@/platform/PlatformContext';
export function ServerTab() {
@@ -11,6 +12,7 @@ export function ServerTab() {
<ConnectionForm />
<ServerStatus />
</div>
<ProviderSettings />
{platform.metadata.isTauri && <UpdateStatus />}
<div className="py-8 text-center text-sm text-muted-foreground">
Created by{' '}
+71
View File
@@ -199,6 +199,77 @@ class ApiClient {
});
}
// Providers
async listProviders(): Promise<{
providers: Array<{
type: string;
name: string;
installed: boolean;
size_mb: number | null;
}>;
installed: string[];
}> {
return this.request('/providers');
}
async getActiveProvider(): Promise<{
provider: string;
health: {
status: string;
provider: string;
version: string | null;
model: string | null;
device: string | null;
};
status: {
model_loaded: boolean;
model_size: string | null;
available_sizes: string[];
gpu_available: boolean | null;
vram_used_mb: number | null;
};
}> {
return this.request('/providers/active');
}
async startProvider(providerType: string): Promise<{
message: string;
provider: {
status: string;
provider: string;
version: string | null;
model: string | null;
device: string | null;
};
}> {
return this.request('/providers/start', {
method: 'POST',
body: JSON.stringify({ provider_type: providerType }),
});
}
async stopProvider(): Promise<{ message: string }> {
return this.request('/providers/stop', {
method: 'POST',
});
}
async downloadProvider(providerType: string): Promise<{
message: string;
provider_type: string;
}> {
return this.request('/providers/download', {
method: 'POST',
body: JSON.stringify({ provider_type: providerType }),
});
}
async deleteProvider(providerType: string): Promise<{ message: string }> {
return this.request(`/providers/${providerType}`, {
method: 'DELETE',
});
}
// History
async listHistory(query?: HistoryQuery): Promise<HistoryListResponse> {
const params = new URLSearchParams();
+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:
+121 -111
View File
@@ -10,14 +10,17 @@
Split the monolithic backend into modular components:
1. **Main App** (~150-200MB): Tauri + FastAPI backend + Whisper + UI/profiles/history
2. **TTS Providers** (downloadable plugins): Separate executables for model inference
1. **Main App**:
- Windows/Linux (~150MB): Tauri + FastAPI backend + Whisper + UI/profiles/history
- macOS (~300MB): Same + MLX bundled for simplicity
2. **TTS Providers** (Windows/Linux only): Downloadable executables for PyTorch CPU/CUDA inference
This architecture solves:
- ✅ GitHub 2GB release artifact limit
- ✅ Frequent app updates without re-downloading large python binaries
- ✅ User choice of compute backend (CPU/GPU/Cloud)
- ✅ Frequent app updates without re-downloading large python binaries (Windows/Linux)
- ✅ User choice of compute backend (CPU/GPU/Cloud) on Windows/Linux
- ✅ Simplified out-of-the-box experience on macOS
- ✅ External provider support (OpenAI, custom servers)
- ✅ Future extensibility
@@ -25,6 +28,7 @@ This architecture solves:
## Architecture Diagram
### Windows / Linux
```
┌─────────────────────────────────────────────────────────┐
│ Voicebox App (Tauri + Backend) ~150MB │
@@ -39,27 +43,43 @@ This architecture solves:
HTTP/IPC │
┌────────────────────────────────┼─────────────────┐
─────────────────┐ ┌─────────────────┐ ┌──────────────────┐
│ TTS Provider: │ │ TTS Provider: │ TTS Provider: │
│ PyTorch CPU │ PyTorch CUDA │ │ MLX (Apple)
│ ~300MB ~2.4GB │ ~800MB
│ Local inference │ │ GPU inference Metal inference │
─────────────────┘ └─────────────────┘ └──────────────────┘
└────────────────────────┴─────────────────────┘
┌─────────────▼──────────────┐
│ Future Providers: │
│ • Remote Server │
│ • OpenAI API │
│ • ElevenLabs │
│ • Custom Docker Container │
└────────────────────────────┘
┌──────────────────────────────────────┐
│ │
▼ ▼
┌─────────────────────┐ ┌─────────────────────┐
│ TTS Provider: │ TTS Provider:
│ PyTorch CPU │ PyTorch CUDA │
│ │ │
│ ~300MB │ ~2.4GB
│ │ │
│ Local inference GPU inference
└─────────────────────┘ └─────────────────────┘
│ │
└──────────────────────────────────────┘
┌─────────────▼──────────────┐
│ Future Providers: │
│ • Remote Server │
│ • OpenAI API │
│ • ElevenLabs │
│ • Custom Docker Container │
└────────────────────────────┘
```
### macOS
```
┌─────────────────────────────────────────────────────────┐
│ Voicebox App (Tauri + Backend) ~300MB │
│ ├─ UI Layer (React) │
│ ├─ Backend (FastAPI) │
│ │ ├─ Voice Profiles │
│ │ ├─ Generation History │
│ │ ├─ Audio Editing / Stories │
│ │ └─ MLX Backend (bundled) │
│ └─ Whisper (bundled, tiny ~50MB) │
│ │
│ No provider downloads needed - works out of the box │
└─────────────────────────────────────────────────────────┘
```
---
@@ -91,18 +111,20 @@ This architecture solves:
#### 1. Main App (voicebox.exe / .app / .AppImage)
**Size:** ~100-150MB
**Windows/Linux Size:** ~100-150MB
**macOS Size:** ~300-350MB (includes MLX)
**Includes:**
- Tauri runtime + React UI
- FastAPI backend (pure Python, no PyTorch)
- FastAPI backend (pure Python, no PyTorch on Windows/Linux)
- Whisper model (tiny, ~50MB)
- SQLite database
- Profile/history/audio editing logic
- Provider management system
- Provider management system (Windows/Linux only)
- **MLX backend (macOS only, bundled)**
**Does NOT include:**
**Does NOT include (Windows/Linux only):**
- PyTorch (CPU or CUDA)
- TTS models (Qwen3-TTS)
@@ -147,23 +169,7 @@ This architecture solves:
---
#### 4. TTS Provider: MLX
**Binary:** `tts-provider-mlx`
**Size:** ~150MB
**Includes:**
- MLX framework
- MLX-optimized Qwen3-TTS
- Metal acceleration
**Platform:** macOS only (Apple Silicon)
**Download source:** Cloudflare R2
---
#### 5. TTS Provider: Remote
#### 4. TTS Provider: Remote
**Binary:** None (built-in config)
**Size:** 0MB
@@ -182,7 +188,7 @@ This architecture solves:
---
#### 6. TTS Provider: OpenAI
#### 5. TTS Provider: OpenAI
**Binary:** None (API wrapper)
**Size:** 0MB
@@ -296,7 +302,10 @@ Model status.
```python
class ProviderManager:
"""Manages TTS provider lifecycle."""
"""Manages TTS provider lifecycle (Windows/Linux only).
Note: macOS uses bundled MLX backend directly, no provider management needed.
"""
def __init__(self):
self.active_provider: Optional[Provider] = None
@@ -308,8 +317,6 @@ class ProviderManager:
return await self._start_local_provider("tts-provider-pytorch-cpu.exe")
elif provider_type == "pytorch-cuda":
return await self._start_local_provider("tts-provider-pytorch-cuda.exe")
elif provider_type == "mlx":
return await self._start_local_provider("tts-provider-mlx")
elif provider_type == "remote":
return self.config["remote_url"]
elif provider_type == "openai":
@@ -434,15 +441,14 @@ class OpenAIProvider(TTSProvider):
```python
class ProviderInstaller:
"""Handles provider download and installation."""
"""Handles provider download and installation (Windows/Linux only)."""
async def download_provider(self, provider_type: str):
"""Download provider binary from R2."""
binary_name = {
"pytorch-cpu": "tts-provider-pytorch-cpu.exe",
"pytorch-cuda": "tts-provider-pytorch-cuda.exe",
"mlx": "tts-provider-mlx"
"pytorch-cuda": "tts-provider-pytorch-cuda.exe"
}[provider_type]
download_url = f"https://downloads.voicebox.sh/providers/v{PROVIDER_VERSION}/{binary_name}"
@@ -525,44 +531,38 @@ export function ProviderSettings() {
)}
</div>
{/* PyTorch CPU */}
<div className="flex items-center justify-between">
<div className="flex items-center space-x-2">
<RadioGroupItem value="pytorch-cpu" id="cpu" />
<Label htmlFor="cpu">
<div className="font-medium">PyTorch CPU</div>
<div className="text-sm text-muted-foreground">
Works on any system, slower inference
</div>
</Label>
</div>
{!installedProviders?.includes("pytorch-cpu") && (
<Button onClick={() => downloadProvider("pytorch-cpu")} size="sm">
Download (300MB)
</Button>
)}
</div>
{/* MLX (macOS only) */}
{isMacOS && (
{/* PyTorch CPU (Windows/Linux only) */}
{!isMacOS && (
<div className="flex items-center justify-between">
<div className="flex items-center space-x-2">
<RadioGroupItem value="mlx" id="mlx" />
<Label htmlFor="mlx">
<div className="font-medium">MLX (Apple Silicon)</div>
<RadioGroupItem value="pytorch-cpu" id="cpu" />
<Label htmlFor="cpu">
<div className="font-medium">PyTorch CPU</div>
<div className="text-sm text-muted-foreground">
Optimized for M1/M2/M3 chips
Works on any system, slower inference
</div>
</Label>
</div>
{!installedProviders?.includes("mlx") && (
<Button onClick={() => downloadProvider("mlx")} size="sm">
Download (800MB)
{!installedProviders?.includes("pytorch-cpu") && (
<Button onClick={() => downloadProvider("pytorch-cpu")} size="sm">
Download (300MB)
</Button>
)}
</div>
)}
{/* MLX bundled (macOS only) */}
{isMacOS && (
<div className="p-3 bg-muted rounded-md">
<div className="text-sm">
<div className="font-medium">MLX (Apple Silicon)</div>
<div className="text-muted-foreground mt-1">
Bundled with the app - optimized for M1/M2/M3 chips
</div>
</div>
</div>
)}
{/* Remote */}
<div className="space-y-2">
<div className="flex items-center space-x-2">
@@ -608,14 +608,18 @@ export function ProviderSettings() {
```
voicebox/
├── backend/
│ ├── main.py # Main FastAPI app (no TTS code)
│ ├── main.py # Main FastAPI app (no TTS on Win/Linux)
│ ├── backends/
│ │ ├── __init__.py # Backend abstraction (existing)
│ │ ├── pytorch_backend.py # PyTorch backend (existing, for reference)
│ │ └── mlx_backend.py # MLX backend (bundled in macOS build only)
│ ├── providers/
│ │ ├── __init__.py # ProviderManager
│ │ ├── base.py # TTSProvider ABC
│ │ ├── __init__.py # ProviderManager (Windows/Linux)
│ │ ├── base.py # TTSProvider Protocol
│ │ ├── local.py # LocalProvider (subprocess)
│ │ ├── remote.py # RemoteProvider (HTTP)
│ │ ├── openai.py # OpenAIProvider (API wrapper)
│ │ └── installer.py # Provider download logic
│ │ └── installer.py # Provider download logic (Windows/Linux)
│ ├── profiles.py # Voice profile management
│ ├── history.py # Generation history
│ ├── transcribe.py # Whisper (still bundled)
@@ -628,27 +632,22 @@ voicebox/
│ │ ├── requirements.txt # torch (CPU), qwen-tts, transformers
│ │ └── build.spec # PyInstaller spec
│ │
── pytorch-cuda/
│ │ ├── main.py # FastAPI server for TTS
│ │ ├── tts_backend.py # PyTorch TTS logic
│ │ ├── requirements.txt # torch+cu121, qwen-tts, transformers
│ │ └── build.spec # PyInstaller spec
│ │
│ └── mlx/
── pytorch-cuda/
│ ├── main.py # FastAPI server for TTS
│ ├── mlx_backend.py # MLX TTS logic
│ ├── requirements.txt # mlx, qwen-tts-mlx
│ ├── tts_backend.py # PyTorch TTS logic
│ ├── requirements.txt # torch+cu121, qwen-tts, transformers
│ └── build.spec # PyInstaller spec
├── app/ # Frontend (Tauri + React)
│ └── src/
│ └── components/
│ └── ServerSettings/
│ └── ProviderSettings.tsx
│ └── ProviderSettings.tsx # Only shown on Windows/Linux
└── tauri/
└── src-tauri/
└── tauri.conf.json # No externalBin for providers
└── tauri.conf.json # No externalBin for providers (Windows/Linux)
# MLX bundled in macOS build
```
---
@@ -671,33 +670,35 @@ voicebox/
### Phase 2: Build Provider Binaries
**Goal:** Create standalone TTS provider executables
**Goal:** Create standalone TTS provider executables (Windows/Linux only)
1. Create separate PyInstaller specs for each provider
2. Build provider executables:
- `tts-provider-pytorch-cpu.exe` (~300MB)
- `tts-provider-pytorch-cuda.exe` (~2.4GB)
- `tts-provider-mlx` (~800MB, macOS)
3. Test subprocess communication
4. Upload providers to Cloudflare R2
**Result:** Provider binaries exist but aren't used yet
**Note:** macOS keeps MLX bundled in main app - no separate provider needed
---
### Phase 3: Remove PyTorch from Main App
**Goal:** Split main app from providers
**Goal:** Split main app from providers (Windows/Linux only)
1. Exclude PyTorch/Qwen3-TTS from main app PyInstaller spec
2. Main app now requires provider download
1. Exclude PyTorch/Qwen3-TTS from Windows/Linux main app PyInstaller spec
2. Windows/Linux app now requires provider download
3. Update GitHub CI to build multiple artifacts:
- `voicebox-{version}-{platform}.exe` (~150MB)
- `voicebox-{version}-windows.exe` (~150MB, no TTS)
- `voicebox-{version}-linux.AppImage` (~150MB, no TTS)
- `voicebox-{version}-macos.app` (~300MB, MLX bundled)
- `tts-provider-pytorch-cpu-{version}.exe`
- `tts-provider-pytorch-cuda-{version}.exe`
- `tts-provider-mlx-{version}` (macOS)
**Result:** Main app is small, providers downloaded separately
**Result:** Windows/Linux apps are small with downloadable providers, macOS app is self-contained
---
@@ -767,7 +768,7 @@ async def check_provider_compatibility(provider_version: str) -> bool:
## User Flows
### First-Time Setup
### First-Time Setup (Windows/Linux)
1. User downloads and installs Voicebox (~150MB)
2. App launches → detects no TTS provider installed
@@ -784,10 +785,6 @@ async def check_provider_compatibility(provider_version: str) -> bool:
✓ Works on any system
✗ Slower inference
[ ] MLX (800MB) [Download]
✓ Fast on Apple Silicon
✗ macOS only (M1/M2/M3)
[ ] Remote Server
URL: ___________________
@@ -799,19 +796,31 @@ async def check_provider_compatibility(provider_version: str) -> bool:
5. Provider installs to AppData/Application Support
6. App starts provider → ready to use
### First-Time Setup (macOS)
1. User downloads and installs Voicebox (~300MB with MLX bundled)
2. App launches → MLX backend is ready immediately
3. No provider setup needed - works out of the box
---
### App Update Flow (No Provider Change)
**Scenario:** Bug fix in UI, no backend changes
**Windows/Linux:**
1. User gets update notification: "Voicebox v0.2.1 available"
2. Downloads update (~150MB, not 2.4GB!)
3. Installs and restarts
4. **Provider stays the same** (no re-download needed)
5. App starts using existing provider
**User experience:** Fast updates, no multi-GB downloads
**macOS:**
1. User gets update notification: "Voicebox v0.2.1 available"
2. Downloads update (~300MB with MLX bundled)
3. Installs and restarts - ready to use
**User experience:** Fast updates, no multi-GB downloads (especially for CUDA users)
---
@@ -846,9 +855,10 @@ async def check_provider_compatibility(provider_version: str) -> bool:
| Benefit | Details |
| ----------------------------- | --------------------------------------------------------- |
| **GitHub Releases Work** | Main app ~150MB << 2GB limit |
| **GitHub Releases Work** | Main app ~150MB (Win/Linux), ~300MB (macOS) << 2GB limit |
| **Fast Updates** | UI/feature updates don't require re-downloading providers |
| **User Choice** | CPU, CUDA, MLX, OpenAI, remote server |
| **User Choice** | CPU, CUDA, OpenAI, remote server (Win/Linux) |
| **macOS Simplicity** | MLX bundled - works immediately, no provider setup needed |
| **External Provider Support** | Users can run their own TTS servers |
| **Bandwidth Savings** | Only download provider once, app updates are small |
| **Future-Proof** | Easy to add new providers (ElevenLabs, custom models) |
+291
View File
@@ -0,0 +1,291 @@
# TTS Provider Architecture
This document explains how Voicebox's modular TTS provider system works.
## Overview
Voicebox uses a **pluggable provider architecture** that separates the main application from TTS inference. This solves several problems:
- **GitHub's 2GB release limit** - CUDA builds are ~2.4GB, too large for GitHub releases
- **Faster app updates** - UI/feature updates don't require re-downloading heavy ML binaries
- **User choice** - Users can pick CPU, CUDA, or external providers based on their hardware
## Architecture Diagram
```
┌─────────────────────────────────────────────────────────────┐
│ Voicebox App │
│ ├─ UI (React) │
│ ├─ Backend (FastAPI) │
│ │ ├─ Voice Profiles │
│ │ ├─ Generation History │
│ │ ├─ Whisper STT (bundled) │
│ │ └─ Provider Manager ◄────────────────┐ │
│ │ │ │
│ └─ providers/ │ │
│ ├─ bundled.py (wraps backends/) │ │
│ └─ local.py (HTTP client)─────────────┼───┐ │
│ │ │ │
└────────────────────────────────────────────┼───┼────────────┘
│ │
┌────────────────────┘ │
│ │ HTTP
▼ ▼
┌──────────────────┐ ┌──────────────────────┐
│ backends/ │ │ Standalone Provider │
│ (bundled on Mac) │ │ (subprocess) │
│ │ │ │
│ - mlx_backend │ │ - FastAPI server │
│ - pytorch_backend│ │ - PyTorch + Qwen-TTS │
└──────────────────┘ │ - Runs on localhost │
└──────────────────────┘
```
## Platform Behavior
| Platform | App Size | TTS Backend | Provider Download |
|----------|----------|-------------|-------------------|
| macOS (Apple Silicon) | ~300MB | MLX bundled | Not needed |
| macOS (Intel) | ~300MB | PyTorch bundled | Not needed |
| Windows | ~150MB | None bundled | Required |
| Linux | ~150MB | None bundled | Required |
### macOS (Apple Silicon)
- MLX backend is **bundled** in the app
- Works immediately after install
- Uses Metal for GPU acceleration
### macOS (Intel)
- PyTorch backend is **bundled** in the app
- Works immediately after install
- Uses CPU inference
### Windows / Linux
- **No TTS bundled** - keeps app small (~150MB)
- On first use, prompts to download a provider
- Provider options:
- **PyTorch CPU** (~300MB) - Works on any system
- **PyTorch CUDA** (~2.4GB) - Fast inference on NVIDIA GPUs
## Directory Structure
```
voicebox/
├── backend/
│ ├── backends/ # Actual TTS implementations
│ │ ├── __init__.py # TTSBackend Protocol
│ │ ├── mlx_backend.py # MLX implementation (macOS)
│ │ └── pytorch_backend.py # PyTorch implementation
│ │
│ └── providers/ # Provider abstraction layer
│ ├── __init__.py # ProviderManager
│ ├── base.py # TTSProvider Protocol
│ ├── bundled.py # Wraps backends/ for bundled use
│ ├── local.py # HTTP client for subprocess providers
│ ├── installer.py # Downloads providers from R2
│ └── types.py # Shared types
└── providers/ # Standalone provider builds
├── pytorch-cpu/
│ ├── main.py # FastAPI server
│ ├── build.py # PyInstaller build script
│ └── requirements.txt
└── pytorch-cuda/
├── main.py # FastAPI server
│ build.py # PyInstaller build script
└── requirements.txt
```
## How Providers Work
### 1. BundledProvider (macOS)
On macOS, the `BundledProvider` directly calls the bundled `backends/` code:
```python
# backend/providers/bundled.py
class BundledProvider:
def __init__(self):
self._backend = get_tts_backend() # MLX or PyTorch
async def generate(self, text, voice_prompt, ...):
return await self._backend.generate(text, voice_prompt, ...)
```
### 2. LocalProvider (Windows/Linux)
On Windows/Linux, the `LocalProvider` communicates with a standalone provider via HTTP:
```python
# backend/providers/local.py
class LocalProvider:
def __init__(self, base_url: str):
self.base_url = base_url # e.g., "http://127.0.0.1:8765"
async def generate(self, text, voice_prompt, ...):
response = await self.client.post(
f"{self.base_url}/tts/generate",
json={"text": text, "voice_prompt": voice_prompt, ...}
)
# Decode audio from response
return audio, sample_rate
```
### 3. Standalone Provider Server
The standalone providers are self-contained FastAPI servers:
```python
# providers/pytorch-cpu/main.py
@app.post("/tts/generate")
async def generate(text: str, voice_prompt: dict, ...):
audio, sr = await backend.generate(text, voice_prompt, ...)
return {"audio": base64_encode(audio), "sample_rate": sr}
```
## Provider API Specification
All providers (local or remote) must implement these HTTP endpoints:
### POST /tts/generate
Generate speech from text.
**Request:**
```json
{
"text": "Hello world!",
"voice_prompt": { /* voice embedding */ },
"language": "en",
"seed": 12345,
"model_size": "1.7B"
}
```
**Response:**
```json
{
"audio": "base64-encoded-wav",
"sample_rate": 24000,
"duration": 2.5
}
```
### POST /tts/create_voice_prompt
Create voice embedding from reference audio.
**Request:** `multipart/form-data`
- `audio`: Audio file
- `reference_text`: Transcript
**Response:**
```json
{
"voice_prompt": { /* voice embedding */ },
"was_cached": false
}
```
### GET /tts/health
Health check.
**Response:**
```json
{
"status": "healthy",
"provider": "pytorch-cuda",
"version": "1.0.0",
"model": "1.7B",
"device": "cuda:0"
}
```
### GET /tts/status
Model status.
**Response:**
```json
{
"model_loaded": true,
"model_size": "1.7B",
"available_sizes": ["0.6B", "1.7B"],
"gpu_available": true,
"vram_used_mb": 1234
}
```
## Provider Lifecycle
### Startup Flow (Windows/Linux)
```
1. App launches
2. ProviderManager checks for installed providers
3. If none installed:
└─ Show setup wizard, prompt download
4. If installed:
├─ Start provider subprocess on random port
├─ Wait for /tts/health to return 200
└─ Create LocalProvider with that URL
5. Generation requests go through LocalProvider → subprocess
```
### Download Flow
```
1. User clicks "Download PyTorch CUDA"
2. Installer downloads from Cloudflare R2:
https://downloads.voicebox.sh/providers/v1.0.0/tts-provider-pytorch-cuda-windows.exe
3. Saved to:
- Windows: %APPDATA%/voicebox/providers/
- Linux: ~/.local/share/voicebox/providers/
4. Provider is now available to start
```
## Building Providers
### Prerequisites
- Python 3.12
- PyInstaller
### Build PyTorch CPU Provider
```bash
cd providers/pytorch-cpu
pip install -r requirements.txt
python build.py
# Output: dist/tts-provider-pytorch-cpu.exe
```
### Build PyTorch CUDA Provider
```bash
cd providers/pytorch-cuda
pip install torch --index-url https://download.pytorch.org/whl/cu121
pip install -r requirements.txt
python build.py
# Output: dist/tts-provider-pytorch-cuda.exe (~2.4GB)
```
## Provider Versioning
Providers have **independent versions** from the app:
- **App version:** `v0.2.0` (frequent updates)
- **Provider version:** `v1.0.0` (rare updates)
Providers only need updates when:
- TTS model changes (new Qwen3-TTS version)
- API spec changes
- Bug fixes in inference code
The app checks provider compatibility on startup.
## Future Providers
The architecture supports additional providers:
- **Remote Server** - Connect to your own TTS server
- **OpenAI API** - Use OpenAI's TTS (requires API key)
- **ElevenLabs** - Cloud TTS service
- **Docker** - Run providers in containers
These would implement the same HTTP API spec.
+82
View File
@@ -0,0 +1,82 @@
"""
PyInstaller build script for PyTorch CPU provider.
"""
import PyInstaller.__main__
import os
import platform
from pathlib import Path
def build_provider():
"""Build PyTorch CPU provider as standalone binary."""
provider_dir = Path(__file__).parent
backend_dir = provider_dir.parent.parent / "backend"
# PyInstaller arguments
args = [
'main.py',
'--onefile',
'--name', 'tts-provider-pytorch-cpu',
]
# Add backend to path
args.extend([
'--paths', str(backend_dir.parent),
])
# Add hidden imports
args.extend([
'--hidden-import', 'backend',
'--hidden-import', 'backend.backends',
'--hidden-import', 'backend.backends.pytorch_backend',
'--hidden-import', 'backend.config',
'--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.tasks',
'--hidden-import', 'torch',
'--hidden-import', 'transformers',
'--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',
'--hidden-import', 'pkg_resources.extern',
'--collect-submodules', 'jaraco',
'--hidden-import', 'fastapi',
'--hidden-import', 'uvicorn',
'--hidden-import', 'soundfile',
'--hidden-import', 'numpy',
'--hidden-import', 'librosa',
])
# Platform-specific extensions
if platform.system() == "Windows":
args[2] = 'tts-provider-pytorch-cpu.exe'
args.extend([
'--noconfirm',
'--clean',
])
# Change to provider directory
os.chdir(provider_dir)
# Run PyInstaller
PyInstaller.__main__.run(args)
binary_name = 'tts-provider-pytorch-cpu'
if platform.system() == "Windows":
binary_name += '.exe'
print(f"Binary built in {provider_dir / 'dist' / binary_name}")
if __name__ == '__main__':
build_provider()
+238
View File
@@ -0,0 +1,238 @@
"""
Standalone TTS provider server for PyTorch CPU.
"""
import argparse
import asyncio
import base64
import io
import sys
from pathlib import Path
from typing import Optional
import numpy as np
import soundfile as sf
from fastapi import FastAPI, File, Form, HTTPException, UploadFile
from fastapi.middleware.cors import CORSMiddleware
import uvicorn
# Add parent directory to path to import backend modules
sys.path.insert(0, str(Path(__file__).parent.parent.parent / "backend"))
from backend.backends.pytorch_backend import PyTorchTTSBackend
app = FastAPI(title="Voicebox TTS Provider - PyTorch CPU")
# CORS middleware
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# Global backend instance
_backend: Optional[PyTorchTTSBackend] = None
def get_backend() -> PyTorchTTSBackend:
"""Get or create backend instance."""
global _backend
if _backend is None:
_backend = PyTorchTTSBackend()
return _backend
@app.get("/tts/health")
async def health():
"""Health check endpoint."""
backend = get_backend()
backend_type = "pytorch-cpu"
model_size = None
if backend.is_loaded():
if hasattr(backend, '_current_model_size') and backend._current_model_size:
model_size = backend._current_model_size
device = backend.device if hasattr(backend, 'device') else "cpu"
return {
"status": "healthy",
"provider": backend_type,
"version": "1.0.0", # TODO: Get from version file
"model": model_size,
"device": device,
}
@app.get("/tts/status")
async def status():
"""Model status endpoint."""
backend = get_backend()
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", "0.6B"]
gpu_available = False
vram_used_mb = None
try:
import torch
gpu_available = torch.cuda.is_available()
if gpu_available:
vram_used_mb = int(torch.cuda.memory_allocated() / 1024 / 1024)
except ImportError:
pass
return {
"model_loaded": backend.is_loaded(),
"model_size": model_size,
"available_sizes": available_sizes,
"gpu_available": gpu_available,
"vram_used_mb": vram_used_mb,
}
@app.post("/tts/generate")
async def generate(
text: str,
voice_prompt: dict,
language: str = "en",
seed: Optional[int] = None,
model_size: str = "1.7B",
):
"""
Generate speech from text.
Request body (JSON):
{
"text": "Hello world!",
"voice_prompt": {...},
"language": "en",
"seed": 12345,
"model_size": "1.7B"
}
"""
backend = get_backend()
# Load model if not loaded or different size
if not backend.is_loaded() or (
hasattr(backend, '_current_model_size') and
backend._current_model_size != model_size
):
await backend.load_model_async(model_size)
# Generate audio
audio, sample_rate = await backend.generate(
text=text,
voice_prompt=voice_prompt,
language=language,
seed=seed,
instruct=None, # TODO: Add instruct support
)
# Convert to base64
buffer = io.BytesIO()
sf.write(buffer, audio, sample_rate, format="WAV")
buffer.seek(0)
audio_bytes = buffer.read()
audio_b64 = base64.b64encode(audio_bytes).decode('utf-8')
# Calculate duration
duration = len(audio) / sample_rate
return {
"audio": audio_b64,
"sample_rate": sample_rate,
"duration": duration,
}
@app.post("/tts/create_voice_prompt")
async def create_voice_prompt(
audio: UploadFile = File(...),
reference_text: str = Form(...),
use_cache: bool = Form(True),
):
"""
Create voice prompt from reference audio.
Request (multipart/form-data):
- audio: Audio file
- reference_text: Transcript
- use_cache: Whether to use cached prompts (default: true)
"""
backend = get_backend()
# Save uploaded file temporarily
import tempfile
with tempfile.NamedTemporaryFile(delete=False, suffix=".wav") as tmp_file:
tmp_path = tmp_file.name
content = await audio.read()
tmp_file.write(content)
try:
# Create voice prompt
voice_prompt, was_cached = await backend.create_voice_prompt(
audio_path=tmp_path,
reference_text=reference_text,
use_cache=use_cache,
)
return {
"voice_prompt": voice_prompt,
"was_cached": was_cached,
}
finally:
# Clean up temp file
Path(tmp_path).unlink(missing_ok=True)
def main():
"""Main entry point."""
parser = argparse.ArgumentParser(description="Voicebox TTS Provider - PyTorch CPU")
parser.add_argument(
"--port",
type=int,
default=0, # 0 means random port
help="Port to bind to",
)
parser.add_argument(
"--data-dir",
type=str,
default=None,
help="Data directory for models and cache",
)
args = parser.parse_args()
# Set data directory if provided
if args.data_dir:
from backend import config
config.set_data_dir(args.data_dir)
# Determine port
port = args.port
if port == 0:
import socket
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.bind(('', 0))
port = s.getsockname()[1]
print(f"Starting TTS Provider (PyTorch CPU) on port {port}")
uvicorn.run(
app,
host="127.0.0.1",
port=port,
log_level="info",
)
if __name__ == "__main__":
main()
+8
View File
@@ -0,0 +1,8 @@
torch>=2.0.0
transformers>=4.30.0
qwen-tts>=0.1.0
fastapi>=0.100.0
uvicorn>=0.23.0
soundfile>=0.12.0
numpy>=1.24.0
librosa>=0.10.0
+84
View File
@@ -0,0 +1,84 @@
"""
PyInstaller build script for PyTorch CUDA provider.
"""
import PyInstaller.__main__
import os
import platform
from pathlib import Path
def build_provider():
"""Build PyTorch CUDA provider as standalone binary."""
provider_dir = Path(__file__).parent
backend_dir = provider_dir.parent.parent / "backend"
# PyInstaller arguments
args = [
'main.py',
'--onefile',
'--name', 'tts-provider-pytorch-cuda',
]
# Add backend to path
args.extend([
'--paths', str(backend_dir.parent),
])
# Add hidden imports
args.extend([
'--hidden-import', 'backend',
'--hidden-import', 'backend.backends',
'--hidden-import', 'backend.backends.pytorch_backend',
'--hidden-import', 'backend.config',
'--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.tasks',
'--hidden-import', 'torch',
'--hidden-import', 'torch.cuda',
'--hidden-import', 'torch.backends.cudnn',
'--hidden-import', 'transformers',
'--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',
'--hidden-import', 'pkg_resources.extern',
'--collect-submodules', 'jaraco',
'--hidden-import', 'fastapi',
'--hidden-import', 'uvicorn',
'--hidden-import', 'soundfile',
'--hidden-import', 'numpy',
'--hidden-import', 'librosa',
])
# Platform-specific extensions
if platform.system() == "Windows":
args[2] = 'tts-provider-pytorch-cuda.exe'
args.extend([
'--noconfirm',
'--clean',
])
# Change to provider directory
os.chdir(provider_dir)
# Run PyInstaller
PyInstaller.__main__.run(args)
binary_name = 'tts-provider-pytorch-cuda'
if platform.system() == "Windows":
binary_name += '.exe'
print(f"Binary built in {provider_dir / 'dist' / binary_name}")
if __name__ == '__main__':
build_provider()
+238
View File
@@ -0,0 +1,238 @@
"""
Standalone TTS provider server for PyTorch CUDA.
"""
import argparse
import asyncio
import base64
import io
import sys
from pathlib import Path
from typing import Optional
import numpy as np
import soundfile as sf
from fastapi import FastAPI, File, Form, HTTPException, UploadFile
from fastapi.middleware.cors import CORSMiddleware
import uvicorn
# Add parent directory to path to import backend modules
sys.path.insert(0, str(Path(__file__).parent.parent.parent / "backend"))
from backend.backends.pytorch_backend import PyTorchTTSBackend
app = FastAPI(title="Voicebox TTS Provider - PyTorch CUDA")
# CORS middleware
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# Global backend instance
_backend: Optional[PyTorchTTSBackend] = None
def get_backend() -> PyTorchTTSBackend:
"""Get or create backend instance."""
global _backend
if _backend is None:
_backend = PyTorchTTSBackend()
return _backend
@app.get("/tts/health")
async def health():
"""Health check endpoint."""
backend = get_backend()
backend_type = "pytorch-cuda"
model_size = None
if backend.is_loaded():
if hasattr(backend, '_current_model_size') and backend._current_model_size:
model_size = backend._current_model_size
device = backend.device if hasattr(backend, 'device') else "cpu"
return {
"status": "healthy",
"provider": backend_type,
"version": "1.0.0", # TODO: Get from version file
"model": model_size,
"device": device,
}
@app.get("/tts/status")
async def status():
"""Model status endpoint."""
backend = get_backend()
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", "0.6B"]
gpu_available = False
vram_used_mb = None
try:
import torch
gpu_available = torch.cuda.is_available()
if gpu_available:
vram_used_mb = int(torch.cuda.memory_allocated() / 1024 / 1024)
except ImportError:
pass
return {
"model_loaded": backend.is_loaded(),
"model_size": model_size,
"available_sizes": available_sizes,
"gpu_available": gpu_available,
"vram_used_mb": vram_used_mb,
}
@app.post("/tts/generate")
async def generate(
text: str,
voice_prompt: dict,
language: str = "en",
seed: Optional[int] = None,
model_size: str = "1.7B",
):
"""
Generate speech from text.
Request body (JSON):
{
"text": "Hello world!",
"voice_prompt": {...},
"language": "en",
"seed": 12345,
"model_size": "1.7B"
}
"""
backend = get_backend()
# Load model if not loaded or different size
if not backend.is_loaded() or (
hasattr(backend, '_current_model_size') and
backend._current_model_size != model_size
):
await backend.load_model_async(model_size)
# Generate audio
audio, sample_rate = await backend.generate(
text=text,
voice_prompt=voice_prompt,
language=language,
seed=seed,
instruct=None, # TODO: Add instruct support
)
# Convert to base64
buffer = io.BytesIO()
sf.write(buffer, audio, sample_rate, format="WAV")
buffer.seek(0)
audio_bytes = buffer.read()
audio_b64 = base64.b64encode(audio_bytes).decode('utf-8')
# Calculate duration
duration = len(audio) / sample_rate
return {
"audio": audio_b64,
"sample_rate": sample_rate,
"duration": duration,
}
@app.post("/tts/create_voice_prompt")
async def create_voice_prompt(
audio: UploadFile = File(...),
reference_text: str = Form(...),
use_cache: bool = Form(True),
):
"""
Create voice prompt from reference audio.
Request (multipart/form-data):
- audio: Audio file
- reference_text: Transcript
- use_cache: Whether to use cached prompts (default: true)
"""
backend = get_backend()
# Save uploaded file temporarily
import tempfile
with tempfile.NamedTemporaryFile(delete=False, suffix=".wav") as tmp_file:
tmp_path = tmp_file.name
content = await audio.read()
tmp_file.write(content)
try:
# Create voice prompt
voice_prompt, was_cached = await backend.create_voice_prompt(
audio_path=tmp_path,
reference_text=reference_text,
use_cache=use_cache,
)
return {
"voice_prompt": voice_prompt,
"was_cached": was_cached,
}
finally:
# Clean up temp file
Path(tmp_path).unlink(missing_ok=True)
def main():
"""Main entry point."""
parser = argparse.ArgumentParser(description="Voicebox TTS Provider - PyTorch CUDA")
parser.add_argument(
"--port",
type=int,
default=0, # 0 means random port
help="Port to bind to",
)
parser.add_argument(
"--data-dir",
type=str,
default=None,
help="Data directory for models and cache",
)
args = parser.parse_args()
# Set data directory if provided
if args.data_dir:
from backend import config
config.set_data_dir(args.data_dir)
# Determine port
port = args.port
if port == 0:
import socket
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.bind(('', 0))
port = s.getsockname()[1]
print(f"Starting TTS Provider (PyTorch CUDA) on port {port}")
uvicorn.run(
app,
host="127.0.0.1",
port=port,
log_level="info",
)
if __name__ == "__main__":
main()
+10
View File
@@ -0,0 +1,10 @@
torch>=2.0.0 --index-url https://download.pytorch.org/whl/cu121
torchvision>=0.15.0 --index-url https://download.pytorch.org/whl/cu121
torchaudio>=2.0.0 --index-url https://download.pytorch.org/whl/cu121
transformers>=4.30.0
qwen-tts>=0.1.0
fastapi>=0.100.0
uvicorn>=0.23.0
soundfile>=0.12.0
numpy>=1.24.0
librosa>=0.10.0