mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 09:05:17 -07:00
Implement TTS provider management system and update release workflow
- Added support for TTS providers in the backend, including endpoints for listing, starting, stopping, and downloading providers. - Enhanced the release workflow to build and upload TTS provider binaries for both Windows and Linux platforms. - Updated the architecture documentation to reflect the new provider system and its benefits for modularity and user experience. - Introduced a new `ProviderSettings` component in the frontend for managing provider configurations.
This commit is contained in:
+132
-12
@@ -6,7 +6,114 @@ on:
|
|||||||
tags:
|
tags:
|
||||||
- "v*"
|
- "v*"
|
||||||
|
|
||||||
|
env:
|
||||||
|
PROVIDER_VERSION: "1.0.0"
|
||||||
|
|
||||||
jobs:
|
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:
|
release:
|
||||||
permissions:
|
permissions:
|
||||||
contents: write
|
contents: write
|
||||||
@@ -14,22 +121,26 @@ jobs:
|
|||||||
fail-fast: false
|
fail-fast: false
|
||||||
matrix:
|
matrix:
|
||||||
include:
|
include:
|
||||||
|
# macOS Apple Silicon - MLX bundled (works out of the box)
|
||||||
- platform: "macos-latest"
|
- platform: "macos-latest"
|
||||||
args: "--target aarch64-apple-darwin"
|
args: "--target aarch64-apple-darwin"
|
||||||
python-version: "3.12"
|
python-version: "3.12"
|
||||||
backend: "mlx"
|
backend: "mlx"
|
||||||
|
# macOS Intel - PyTorch bundled (smaller user base, keep simple)
|
||||||
- platform: "macos-15-intel"
|
- platform: "macos-15-intel"
|
||||||
args: "--target x86_64-apple-darwin"
|
args: "--target x86_64-apple-darwin"
|
||||||
python-version: "3.12"
|
python-version: "3.12"
|
||||||
backend: "pytorch"
|
backend: "pytorch"
|
||||||
|
# Linux - No TTS bundled, providers downloaded separately
|
||||||
# - platform: 'ubuntu-22.04'
|
# - platform: 'ubuntu-22.04'
|
||||||
# args: ''
|
# args: ''
|
||||||
# python-version: '3.12'
|
# python-version: '3.12'
|
||||||
# backend: 'pytorch'
|
# backend: 'none'
|
||||||
|
# Windows - No TTS bundled, providers downloaded separately
|
||||||
- platform: "windows-latest"
|
- platform: "windows-latest"
|
||||||
args: ""
|
args: ""
|
||||||
python-version: "3.12"
|
python-version: "3.12"
|
||||||
backend: "pytorch"
|
backend: "none"
|
||||||
|
|
||||||
runs-on: ${{ matrix.platform }}
|
runs-on: ${{ matrix.platform }}
|
||||||
|
|
||||||
@@ -55,23 +166,27 @@ jobs:
|
|||||||
python-version: ${{ matrix.python-version }}
|
python-version: ${{ matrix.python-version }}
|
||||||
cache: "pip"
|
cache: "pip"
|
||||||
|
|
||||||
- name: Install Python dependencies
|
- name: Install Python dependencies (with TTS)
|
||||||
|
if: matrix.backend != 'none'
|
||||||
run: |
|
run: |
|
||||||
python -m pip install --upgrade pip
|
python -m pip install --upgrade pip
|
||||||
pip install pyinstaller
|
pip install pyinstaller
|
||||||
pip install -r backend/requirements.txt
|
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)
|
- name: Install MLX dependencies (Apple Silicon only)
|
||||||
if: matrix.backend == 'mlx'
|
if: matrix.backend == 'mlx'
|
||||||
run: |
|
run: |
|
||||||
pip install -r backend/requirements-mlx.txt
|
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)
|
- name: Build Python server (Linux/macOS)
|
||||||
if: matrix.platform != 'windows-latest'
|
if: matrix.platform != 'windows-latest'
|
||||||
run: |
|
run: |
|
||||||
@@ -148,10 +263,15 @@ jobs:
|
|||||||
See the assets below to download and install this version.
|
See the assets below to download and install this version.
|
||||||
|
|
||||||
### Installation
|
### 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
|
- **macOS (Intel)**: Download the `x64.dmg` file - uses PyTorch
|
||||||
- **Windows**: Download the `.msi` installer
|
- **Windows**: Download the `.msi` installer - requires downloading a TTS provider on first use
|
||||||
- **Linux**: Download the `.AppImage` or `.deb` package
|
- **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.
|
The app includes automatic updates - future updates will be installed automatically.
|
||||||
releaseDraft: true
|
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 { ConnectionForm } from '@/components/ServerSettings/ConnectionForm';
|
||||||
import { ServerStatus } from '@/components/ServerSettings/ServerStatus';
|
import { ServerStatus } from '@/components/ServerSettings/ServerStatus';
|
||||||
import { UpdateStatus } from '@/components/ServerSettings/UpdateStatus';
|
import { UpdateStatus } from '@/components/ServerSettings/UpdateStatus';
|
||||||
|
import { ProviderSettings } from '@/components/ServerSettings/ProviderSettings';
|
||||||
import { usePlatform } from '@/platform/PlatformContext';
|
import { usePlatform } from '@/platform/PlatformContext';
|
||||||
|
|
||||||
export function ServerTab() {
|
export function ServerTab() {
|
||||||
@@ -11,6 +12,7 @@ export function ServerTab() {
|
|||||||
<ConnectionForm />
|
<ConnectionForm />
|
||||||
<ServerStatus />
|
<ServerStatus />
|
||||||
</div>
|
</div>
|
||||||
|
<ProviderSettings />
|
||||||
{platform.metadata.isTauri && <UpdateStatus />}
|
{platform.metadata.isTauri && <UpdateStatus />}
|
||||||
<div className="py-8 text-center text-sm text-muted-foreground">
|
<div className="py-8 text-center text-sm text-muted-foreground">
|
||||||
Created by{' '}
|
Created by{' '}
|
||||||
|
|||||||
@@ -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
|
// History
|
||||||
async listHistory(query?: HistoryQuery): Promise<HistoryListResponse> {
|
async listHistory(query?: HistoryQuery): Promise<HistoryListResponse> {
|
||||||
const params = new URLSearchParams();
|
const params = new URLSearchParams();
|
||||||
|
|||||||
+15
-17
@@ -30,7 +30,7 @@ def build_server():
|
|||||||
args.extend(['--paths', str(qwen_tts_path)])
|
args.extend(['--paths', str(qwen_tts_path)])
|
||||||
print(f"Using local qwen_tts source from: {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([
|
args.extend([
|
||||||
'--hidden-import', 'backend',
|
'--hidden-import', 'backend',
|
||||||
'--hidden-import', 'backend.main',
|
'--hidden-import', 'backend.main',
|
||||||
@@ -42,38 +42,30 @@ def build_server():
|
|||||||
'--hidden-import', 'backend.tts',
|
'--hidden-import', 'backend.tts',
|
||||||
'--hidden-import', 'backend.transcribe',
|
'--hidden-import', 'backend.transcribe',
|
||||||
'--hidden-import', 'backend.platform_detect',
|
'--hidden-import', 'backend.platform_detect',
|
||||||
'--hidden-import', 'backend.backends',
|
'--hidden-import', 'backend.providers',
|
||||||
'--hidden-import', 'backend.backends.pytorch_backend',
|
'--hidden-import', 'backend.providers.base',
|
||||||
|
'--hidden-import', 'backend.providers.bundled',
|
||||||
|
'--hidden-import', 'backend.providers.types',
|
||||||
'--hidden-import', 'backend.utils.audio',
|
'--hidden-import', 'backend.utils.audio',
|
||||||
'--hidden-import', 'backend.utils.cache',
|
'--hidden-import', 'backend.utils.cache',
|
||||||
'--hidden-import', 'backend.utils.progress',
|
'--hidden-import', 'backend.utils.progress',
|
||||||
'--hidden-import', 'backend.utils.hf_progress',
|
'--hidden-import', 'backend.utils.hf_progress',
|
||||||
'--hidden-import', 'backend.utils.validation',
|
'--hidden-import', 'backend.utils.validation',
|
||||||
'--hidden-import', 'torch',
|
|
||||||
'--hidden-import', 'transformers',
|
|
||||||
'--hidden-import', 'fastapi',
|
'--hidden-import', 'fastapi',
|
||||||
'--hidden-import', 'uvicorn',
|
'--hidden-import', 'uvicorn',
|
||||||
'--hidden-import', 'sqlalchemy',
|
'--hidden-import', 'sqlalchemy',
|
||||||
'--hidden-import', 'librosa',
|
'--hidden-import', 'librosa',
|
||||||
'--hidden-import', 'soundfile',
|
'--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
|
# Fix for pkg_resources and jaraco namespace packages
|
||||||
'--hidden-import', 'pkg_resources.extern',
|
'--hidden-import', 'pkg_resources.extern',
|
||||||
'--collect-submodules', 'jaraco',
|
'--collect-submodules', 'jaraco',
|
||||||
])
|
])
|
||||||
|
|
||||||
# Add MLX-specific imports if building on Apple Silicon
|
# Platform-specific TTS backend handling
|
||||||
if is_apple_silicon():
|
if is_apple_silicon():
|
||||||
print("Building for Apple Silicon - including MLX dependencies")
|
print("Building for Apple Silicon - including MLX dependencies (bundled)")
|
||||||
args.extend([
|
args.extend([
|
||||||
|
'--hidden-import', 'backend.backends',
|
||||||
'--hidden-import', 'backend.backends.mlx_backend',
|
'--hidden-import', 'backend.backends.mlx_backend',
|
||||||
'--hidden-import', 'mlx',
|
'--hidden-import', 'mlx',
|
||||||
'--hidden-import', 'mlx.core',
|
'--hidden-import', 'mlx.core',
|
||||||
@@ -88,7 +80,13 @@ def build_server():
|
|||||||
'--collect-data', 'mlx_audio',
|
'--collect-data', 'mlx_audio',
|
||||||
])
|
])
|
||||||
else:
|
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([
|
args.extend([
|
||||||
'--noconfirm',
|
'--noconfirm',
|
||||||
|
|||||||
+200
-15
@@ -29,6 +29,8 @@ from .utils.progress import get_progress_manager
|
|||||||
from .utils.tasks import get_task_manager
|
from .utils.tasks import get_task_manager
|
||||||
from .utils.cache import clear_voice_prompt_cache
|
from .utils.cache import clear_voice_prompt_cache
|
||||||
from .platform_detect import get_backend_type
|
from .platform_detect import get_backend_type
|
||||||
|
from .providers import get_provider_manager
|
||||||
|
from .providers.types import ProviderType
|
||||||
|
|
||||||
app = FastAPI(
|
app = FastAPI(
|
||||||
title="voicebox API",
|
title="voicebox API",
|
||||||
@@ -74,7 +76,7 @@ async def health():
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
import os
|
import os
|
||||||
|
|
||||||
tts_model = tts.get_tts_model()
|
tts_model = await tts.get_tts_model_async()
|
||||||
backend_type = get_backend_type()
|
backend_type = get_backend_type()
|
||||||
|
|
||||||
# Check for GPU availability (CUDA or MPS)
|
# Check for GPU availability (CUDA or MPS)
|
||||||
@@ -549,7 +551,7 @@ async def generate_speech(
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Generate audio
|
# 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)
|
# Load the requested model size if different from current (async to not block)
|
||||||
model_size = data.model_size or "1.7B"
|
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"):
|
async def load_model(model_size: str = "1.7B"):
|
||||||
"""Manually load TTS model."""
|
"""Manually load TTS model."""
|
||||||
try:
|
try:
|
||||||
tts_model = tts.get_tts_model()
|
tts_model = await tts.get_tts_model_async()
|
||||||
await tts_model.load_model_async(model_size)
|
await tts_model.load_model(model_size)
|
||||||
return {"message": f"Model {model_size} loaded successfully"}
|
return {"message": f"Model {model_size} loaded successfully"}
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
raise HTTPException(status_code=500, detail=str(e))
|
raise HTTPException(status_code=500, detail=str(e))
|
||||||
@@ -1172,10 +1174,10 @@ async def get_model_status():
|
|||||||
except ImportError:
|
except ImportError:
|
||||||
use_scan_cache = False
|
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."""
|
"""Check if TTS model is loaded with specific size."""
|
||||||
try:
|
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
|
return tts_model.is_loaded() and getattr(tts_model, 'model_size', None) == model_size
|
||||||
except Exception:
|
except Exception:
|
||||||
return False
|
return False
|
||||||
@@ -1211,14 +1213,14 @@ async def get_model_status():
|
|||||||
"display_name": "Qwen TTS 1.7B",
|
"display_name": "Qwen TTS 1.7B",
|
||||||
"hf_repo_id": tts_1_7b_id,
|
"hf_repo_id": tts_1_7b_id,
|
||||||
"model_size": "1.7B",
|
"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",
|
"model_name": "qwen-tts-0.6B",
|
||||||
"display_name": "Qwen TTS 0.6B",
|
"display_name": "Qwen TTS 0.6B",
|
||||||
"hf_repo_id": tts_0_6b_id,
|
"hf_repo_id": tts_0_6b_id,
|
||||||
"model_size": "0.6B",
|
"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",
|
"model_name": "whisper-base",
|
||||||
@@ -1356,7 +1358,11 @@ async def get_model_status():
|
|||||||
|
|
||||||
# Check if loaded in memory
|
# Check if loaded in memory
|
||||||
try:
|
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:
|
except Exception:
|
||||||
loaded = False
|
loaded = False
|
||||||
|
|
||||||
@@ -1379,7 +1385,11 @@ async def get_model_status():
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
# If check fails, try to at least check if loaded
|
# If check fails, try to at least check if loaded
|
||||||
try:
|
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:
|
except Exception:
|
||||||
loaded = False
|
loaded = False
|
||||||
|
|
||||||
@@ -1406,14 +1416,24 @@ async def trigger_model_download(request: models.ModelDownloadRequest):
|
|||||||
task_manager = get_task_manager()
|
task_manager = get_task_manager()
|
||||||
progress_manager = get_progress_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 = {
|
model_configs = {
|
||||||
"qwen-tts-1.7B": {
|
"qwen-tts-1.7B": {
|
||||||
"model_size": "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": {
|
"qwen-tts-0.6B": {
|
||||||
"model_size": "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": {
|
"whisper-base": {
|
||||||
"model_size": "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"}
|
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}")
|
@app.delete("/models/{model_name}")
|
||||||
async def delete_model(model_name: str):
|
async def delete_model(model_name: str):
|
||||||
"""Delete a downloaded model from the HuggingFace cache."""
|
"""Delete a downloaded model from the HuggingFace cache."""
|
||||||
@@ -1522,9 +1707,9 @@ async def delete_model(model_name: str):
|
|||||||
try:
|
try:
|
||||||
# Check if model is loaded and unload it first
|
# Check if model is loaded and unload it first
|
||||||
if config["model_type"] == "tts":
|
if config["model_type"] == "tts":
|
||||||
tts_model = tts.get_tts_model()
|
tts_model = await tts.get_tts_model_async()
|
||||||
if tts_model.is_loaded() and tts_model.model_size == config["model_size"]:
|
if tts_model.is_loaded() and getattr(tts_model, 'model_size', None) == config["model_size"]:
|
||||||
tts.unload_tts_model()
|
tts_model.unload_model()
|
||||||
elif config["model_type"] == "whisper":
|
elif config["model_type"] == "whisper":
|
||||||
whisper_model = transcribe.get_whisper_model()
|
whisper_model = transcribe.get_whisper_model()
|
||||||
if whisper_model.is_loaded() and whisper_model.model_size == config["model_size"]:
|
if whisper_model.is_loaded() and whisper_model.model_size == config["model_size"]:
|
||||||
|
|||||||
@@ -0,0 +1,220 @@
|
|||||||
|
"""
|
||||||
|
Provider management system for TTS providers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Optional
|
||||||
|
import platform
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from .base import TTSProvider
|
||||||
|
from .types import ProviderType
|
||||||
|
from .bundled import BundledProvider
|
||||||
|
from .local import LocalProvider
|
||||||
|
from .installer import get_provider_binary_path
|
||||||
|
from ..config import get_data_dir
|
||||||
|
import subprocess
|
||||||
|
import socket
|
||||||
|
|
||||||
|
|
||||||
|
class ProviderManager:
|
||||||
|
"""Manages TTS provider lifecycle."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self.active_provider: Optional[TTSProvider] = None
|
||||||
|
self._default_provider: Optional[TTSProvider] = None
|
||||||
|
self._provider_process: Optional[subprocess.Popen] = None
|
||||||
|
self._provider_port: Optional[int] = None
|
||||||
|
|
||||||
|
def _get_default_provider(self) -> TTSProvider:
|
||||||
|
"""Get the default bundled provider."""
|
||||||
|
if self._default_provider is None:
|
||||||
|
self._default_provider = BundledProvider()
|
||||||
|
return self._default_provider
|
||||||
|
|
||||||
|
async def get_active_provider(self) -> TTSProvider:
|
||||||
|
"""
|
||||||
|
Get the currently active provider.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Active TTS provider instance
|
||||||
|
"""
|
||||||
|
if self.active_provider is None:
|
||||||
|
# Default to bundled provider
|
||||||
|
self.active_provider = self._get_default_provider()
|
||||||
|
return self.active_provider
|
||||||
|
|
||||||
|
async def start_provider(self, provider_type: str) -> None:
|
||||||
|
"""
|
||||||
|
Start a TTS provider.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
provider_type: Type of provider to start
|
||||||
|
"""
|
||||||
|
if provider_type in ["bundled-mlx", "bundled-pytorch"]:
|
||||||
|
# Use bundled provider
|
||||||
|
self.active_provider = self._get_default_provider()
|
||||||
|
elif provider_type in ["pytorch-cpu", "pytorch-cuda"]:
|
||||||
|
# Start local provider subprocess
|
||||||
|
provider_path = get_provider_binary_path(provider_type)
|
||||||
|
if not provider_path or not provider_path.exists():
|
||||||
|
raise ValueError(f"Provider {provider_type} is not installed. Please download it first.")
|
||||||
|
|
||||||
|
# Find a free port
|
||||||
|
port = self._get_free_port()
|
||||||
|
|
||||||
|
# Start provider subprocess
|
||||||
|
from ..config import get_data_dir
|
||||||
|
process = subprocess.Popen(
|
||||||
|
[
|
||||||
|
str(provider_path),
|
||||||
|
"--port", str(port),
|
||||||
|
"--data-dir", str(get_data_dir()),
|
||||||
|
],
|
||||||
|
stdout=subprocess.PIPE,
|
||||||
|
stderr=subprocess.PIPE,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Wait for provider to be ready
|
||||||
|
base_url = f"http://127.0.0.1:{port}"
|
||||||
|
await self._wait_for_provider_health(base_url, timeout=30)
|
||||||
|
|
||||||
|
# Create LocalProvider instance
|
||||||
|
self.active_provider = LocalProvider(base_url)
|
||||||
|
self._provider_process = process
|
||||||
|
self._provider_port = port
|
||||||
|
elif provider_type == "remote":
|
||||||
|
# Remote provider - will be implemented in Phase 5
|
||||||
|
raise NotImplementedError("Remote provider not yet implemented")
|
||||||
|
elif provider_type == "openai":
|
||||||
|
# OpenAI provider - will be implemented in Phase 5
|
||||||
|
raise NotImplementedError("OpenAI provider not yet implemented")
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unknown provider type: {provider_type}")
|
||||||
|
|
||||||
|
async def stop_provider(self) -> None:
|
||||||
|
"""Stop the active provider."""
|
||||||
|
if self.active_provider:
|
||||||
|
# Only stop if it's not the default bundled provider
|
||||||
|
if self.active_provider is not self._default_provider:
|
||||||
|
if hasattr(self.active_provider, 'stop'):
|
||||||
|
await self.active_provider.stop()
|
||||||
|
self.active_provider = None
|
||||||
|
|
||||||
|
# Stop subprocess if running
|
||||||
|
if self._provider_process:
|
||||||
|
self._provider_process.terminate()
|
||||||
|
try:
|
||||||
|
self._provider_process.wait(timeout=5)
|
||||||
|
except subprocess.TimeoutExpired:
|
||||||
|
self._provider_process.kill()
|
||||||
|
self._provider_process = None
|
||||||
|
self._provider_port = None
|
||||||
|
|
||||||
|
async def list_installed(self) -> list[str]:
|
||||||
|
"""
|
||||||
|
List installed provider types.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List of installed provider type strings
|
||||||
|
"""
|
||||||
|
installed = []
|
||||||
|
|
||||||
|
# Bundled providers are always available
|
||||||
|
system = platform.system()
|
||||||
|
machine = platform.machine()
|
||||||
|
|
||||||
|
if system == "Darwin" and machine == "arm64":
|
||||||
|
installed.append("bundled-mlx")
|
||||||
|
else:
|
||||||
|
installed.append("bundled-pytorch")
|
||||||
|
|
||||||
|
# Check for downloaded providers (Phase 2)
|
||||||
|
providers_dir = _get_providers_dir()
|
||||||
|
if providers_dir.exists():
|
||||||
|
for provider_file in providers_dir.glob("tts-provider-*"):
|
||||||
|
if provider_file.is_file() and provider_file.stat().st_size > 0:
|
||||||
|
name = provider_file.name
|
||||||
|
if "pytorch-cpu" in name:
|
||||||
|
installed.append("pytorch-cpu")
|
||||||
|
elif "pytorch-cuda" in name:
|
||||||
|
installed.append("pytorch-cuda")
|
||||||
|
|
||||||
|
return installed
|
||||||
|
|
||||||
|
async def get_provider_info(self, provider_type: str) -> dict:
|
||||||
|
"""
|
||||||
|
Get information about a provider.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
provider_type: Type of provider
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Provider information dictionary
|
||||||
|
"""
|
||||||
|
if provider_type in ["bundled-mlx", "bundled-pytorch"]:
|
||||||
|
return {
|
||||||
|
"type": provider_type,
|
||||||
|
"name": "Bundled Provider",
|
||||||
|
"installed": True,
|
||||||
|
"size_mb": None, # Bundled, no separate size
|
||||||
|
}
|
||||||
|
elif provider_type == "pytorch-cpu":
|
||||||
|
return {
|
||||||
|
"type": provider_type,
|
||||||
|
"name": "PyTorch CPU",
|
||||||
|
"installed": provider_type in await self.list_installed(),
|
||||||
|
"size_mb": 300,
|
||||||
|
}
|
||||||
|
elif provider_type == "pytorch-cuda":
|
||||||
|
return {
|
||||||
|
"type": provider_type,
|
||||||
|
"name": "PyTorch CUDA",
|
||||||
|
"installed": provider_type in await self.list_installed(),
|
||||||
|
"size_mb": 2400,
|
||||||
|
}
|
||||||
|
else:
|
||||||
|
return {
|
||||||
|
"type": provider_type,
|
||||||
|
"name": provider_type,
|
||||||
|
"installed": False,
|
||||||
|
"size_mb": None,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _get_free_port(self) -> int:
|
||||||
|
"""Get a free port for the provider server."""
|
||||||
|
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||||
|
s.bind(('', 0))
|
||||||
|
return s.getsockname()[1]
|
||||||
|
|
||||||
|
async def _wait_for_provider_health(self, base_url: str, timeout: int = 30) -> None:
|
||||||
|
"""Wait for provider to become healthy."""
|
||||||
|
import httpx
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
start_time = asyncio.get_event_loop().time()
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
async with httpx.AsyncClient(timeout=2.0) as client:
|
||||||
|
response = await client.get(f"{base_url}/tts/health")
|
||||||
|
if response.status_code == 200:
|
||||||
|
return
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if asyncio.get_event_loop().time() - start_time > timeout:
|
||||||
|
raise TimeoutError(f"Provider did not become healthy within {timeout} seconds")
|
||||||
|
|
||||||
|
await asyncio.sleep(0.5)
|
||||||
|
|
||||||
|
|
||||||
|
# Global provider manager instance
|
||||||
|
_provider_manager: Optional[ProviderManager] = None
|
||||||
|
|
||||||
|
|
||||||
|
def get_provider_manager() -> ProviderManager:
|
||||||
|
"""Get the global provider manager instance."""
|
||||||
|
global _provider_manager
|
||||||
|
if _provider_manager is None:
|
||||||
|
_provider_manager = ProviderManager()
|
||||||
|
return _provider_manager
|
||||||
@@ -0,0 +1,97 @@
|
|||||||
|
"""
|
||||||
|
Base protocol for TTS providers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Protocol, Optional, Tuple
|
||||||
|
from typing_extensions import runtime_checkable
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
from .types import ProviderHealth, ProviderStatus
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class TTSProvider(Protocol):
|
||||||
|
"""Protocol for TTS provider implementations."""
|
||||||
|
|
||||||
|
async def generate(
|
||||||
|
self,
|
||||||
|
text: str,
|
||||||
|
voice_prompt: dict,
|
||||||
|
language: str = "en",
|
||||||
|
seed: Optional[int] = None,
|
||||||
|
instruct: Optional[str] = None,
|
||||||
|
) -> Tuple[np.ndarray, int]:
|
||||||
|
"""
|
||||||
|
Generate speech audio from text.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text: Text to synthesize
|
||||||
|
voice_prompt: Voice prompt dictionary
|
||||||
|
language: Language code
|
||||||
|
seed: Random seed for reproducibility
|
||||||
|
instruct: Delivery instructions
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (audio_array, sample_rate)
|
||||||
|
"""
|
||||||
|
...
|
||||||
|
|
||||||
|
async def create_voice_prompt(
|
||||||
|
self,
|
||||||
|
audio_path: str,
|
||||||
|
reference_text: str,
|
||||||
|
use_cache: bool = True,
|
||||||
|
) -> Tuple[dict, bool]:
|
||||||
|
"""
|
||||||
|
Create voice prompt from reference audio.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
audio_path: Path to reference audio file
|
||||||
|
reference_text: Transcript of the audio
|
||||||
|
use_cache: Whether to use cached prompts
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (voice_prompt_dict, was_cached)
|
||||||
|
"""
|
||||||
|
...
|
||||||
|
|
||||||
|
async def combine_voice_prompts(
|
||||||
|
self,
|
||||||
|
audio_paths: list[str],
|
||||||
|
reference_texts: list[str],
|
||||||
|
) -> Tuple[np.ndarray, str]:
|
||||||
|
"""
|
||||||
|
Combine multiple voice prompts.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
audio_paths: List of audio file paths
|
||||||
|
reference_texts: List of reference texts
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (combined_audio_array, combined_text)
|
||||||
|
"""
|
||||||
|
...
|
||||||
|
|
||||||
|
async def load_model(self, model_size: str) -> None:
|
||||||
|
"""Load TTS model."""
|
||||||
|
...
|
||||||
|
|
||||||
|
def unload_model(self) -> None:
|
||||||
|
"""Unload model to free memory."""
|
||||||
|
...
|
||||||
|
|
||||||
|
def is_loaded(self) -> bool:
|
||||||
|
"""Check if model is loaded."""
|
||||||
|
...
|
||||||
|
|
||||||
|
def _get_model_path(self, model_size: str) -> str:
|
||||||
|
"""Get model path for a given size."""
|
||||||
|
...
|
||||||
|
|
||||||
|
async def health(self) -> ProviderHealth:
|
||||||
|
"""Get provider health status."""
|
||||||
|
...
|
||||||
|
|
||||||
|
async def status(self) -> ProviderStatus:
|
||||||
|
"""Get provider model status."""
|
||||||
|
...
|
||||||
@@ -0,0 +1,139 @@
|
|||||||
|
"""
|
||||||
|
Bundled provider that wraps existing MLX/PyTorch backends.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Optional, Tuple
|
||||||
|
import numpy as np
|
||||||
|
import platform
|
||||||
|
|
||||||
|
from .base import TTSProvider
|
||||||
|
from .types import ProviderHealth, ProviderStatus
|
||||||
|
from ..backends import get_tts_backend, TTSBackend
|
||||||
|
from ..platform_detect import get_backend_type
|
||||||
|
|
||||||
|
|
||||||
|
class BundledProvider:
|
||||||
|
"""Provider that wraps the existing bundled TTS backend."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._backend: Optional[TTSBackend] = None
|
||||||
|
|
||||||
|
def _get_backend(self) -> TTSBackend:
|
||||||
|
"""Get or create backend instance."""
|
||||||
|
if self._backend is None:
|
||||||
|
self._backend = get_tts_backend()
|
||||||
|
return self._backend
|
||||||
|
|
||||||
|
async def generate(
|
||||||
|
self,
|
||||||
|
text: str,
|
||||||
|
voice_prompt: dict,
|
||||||
|
language: str = "en",
|
||||||
|
seed: Optional[int] = None,
|
||||||
|
instruct: Optional[str] = None,
|
||||||
|
) -> Tuple[np.ndarray, int]:
|
||||||
|
"""Generate speech audio."""
|
||||||
|
backend = self._get_backend()
|
||||||
|
return await backend.generate(text, voice_prompt, language, seed, instruct)
|
||||||
|
|
||||||
|
async def create_voice_prompt(
|
||||||
|
self,
|
||||||
|
audio_path: str,
|
||||||
|
reference_text: str,
|
||||||
|
use_cache: bool = True,
|
||||||
|
) -> Tuple[dict, bool]:
|
||||||
|
"""Create voice prompt from reference audio."""
|
||||||
|
backend = self._get_backend()
|
||||||
|
return await backend.create_voice_prompt(audio_path, reference_text, use_cache)
|
||||||
|
|
||||||
|
async def combine_voice_prompts(
|
||||||
|
self,
|
||||||
|
audio_paths: list[str],
|
||||||
|
reference_texts: list[str],
|
||||||
|
) -> Tuple[np.ndarray, str]:
|
||||||
|
"""Combine multiple voice prompts."""
|
||||||
|
backend = self._get_backend()
|
||||||
|
return await backend.combine_voice_prompts(audio_paths, reference_texts)
|
||||||
|
|
||||||
|
async def load_model(self, model_size: str) -> None:
|
||||||
|
"""Load TTS model."""
|
||||||
|
backend = self._get_backend()
|
||||||
|
# Backends use load_model_async, but Protocol defines load_model
|
||||||
|
if hasattr(backend, 'load_model_async'):
|
||||||
|
await backend.load_model_async(model_size)
|
||||||
|
else:
|
||||||
|
await backend.load_model(model_size)
|
||||||
|
|
||||||
|
def unload_model(self) -> None:
|
||||||
|
"""Unload model to free memory."""
|
||||||
|
backend = self._get_backend()
|
||||||
|
backend.unload_model()
|
||||||
|
|
||||||
|
def is_loaded(self) -> bool:
|
||||||
|
"""Check if model is loaded."""
|
||||||
|
backend = self._get_backend()
|
||||||
|
return backend.is_loaded()
|
||||||
|
|
||||||
|
def _get_model_path(self, model_size: str) -> str:
|
||||||
|
"""Get model path for a given size."""
|
||||||
|
backend = self._get_backend()
|
||||||
|
return backend._get_model_path(model_size)
|
||||||
|
|
||||||
|
async def health(self) -> ProviderHealth:
|
||||||
|
"""Get provider health status."""
|
||||||
|
backend = self._get_backend()
|
||||||
|
backend_type = get_backend_type()
|
||||||
|
|
||||||
|
model_size = None
|
||||||
|
if backend.is_loaded():
|
||||||
|
# Try to get current model size from backend
|
||||||
|
if hasattr(backend, '_current_model_size') and backend._current_model_size:
|
||||||
|
model_size = backend._current_model_size
|
||||||
|
|
||||||
|
device = None
|
||||||
|
if backend_type == "mlx":
|
||||||
|
device = "metal"
|
||||||
|
elif hasattr(backend, 'device'):
|
||||||
|
device = backend.device
|
||||||
|
|
||||||
|
return ProviderHealth(
|
||||||
|
status="healthy",
|
||||||
|
provider=f"bundled-{backend_type}",
|
||||||
|
version=None, # Provider versioning not implemented yet
|
||||||
|
model=model_size,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def status(self) -> ProviderStatus:
|
||||||
|
"""Get provider model status."""
|
||||||
|
backend = self._get_backend()
|
||||||
|
backend_type = get_backend_type()
|
||||||
|
|
||||||
|
model_size = None
|
||||||
|
if backend.is_loaded():
|
||||||
|
if hasattr(backend, '_current_model_size') and backend._current_model_size:
|
||||||
|
model_size = backend._current_model_size
|
||||||
|
|
||||||
|
available_sizes = ["1.7B"]
|
||||||
|
if backend_type == "pytorch":
|
||||||
|
available_sizes.append("0.6B")
|
||||||
|
|
||||||
|
gpu_available = None
|
||||||
|
vram_used_mb = None
|
||||||
|
|
||||||
|
if backend_type == "pytorch":
|
||||||
|
try:
|
||||||
|
import torch
|
||||||
|
gpu_available = torch.cuda.is_available()
|
||||||
|
if gpu_available:
|
||||||
|
vram_used_mb = torch.cuda.memory_allocated() / 1024 / 1024
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
return ProviderStatus(
|
||||||
|
model_loaded=backend.is_loaded(),
|
||||||
|
model_size=model_size,
|
||||||
|
available_sizes=available_sizes,
|
||||||
|
gpu_available=gpu_available,
|
||||||
|
vram_used_mb=int(vram_used_mb) if vram_used_mb else None,
|
||||||
|
)
|
||||||
@@ -0,0 +1,211 @@
|
|||||||
|
"""
|
||||||
|
Provider download and installation manager.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import httpx
|
||||||
|
import platform
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
from .types import ProviderType
|
||||||
|
from ..utils.progress import get_progress_manager
|
||||||
|
from ..utils.tasks import get_task_manager
|
||||||
|
|
||||||
|
|
||||||
|
# Provider version (independent of app version)
|
||||||
|
PROVIDER_VERSION = "1.0.0"
|
||||||
|
|
||||||
|
# Base URL for provider downloads (Cloudflare R2)
|
||||||
|
PROVIDER_DOWNLOAD_BASE_URL = "https://downloads.voicebox.sh/providers"
|
||||||
|
|
||||||
|
|
||||||
|
def _get_providers_dir() -> Path:
|
||||||
|
"""Get the directory where providers are stored."""
|
||||||
|
system = platform.system()
|
||||||
|
|
||||||
|
if system == "Windows":
|
||||||
|
appdata = Path.home() / "AppData" / "Roaming"
|
||||||
|
elif system == "Darwin":
|
||||||
|
appdata = Path.home() / "Library" / "Application Support"
|
||||||
|
else: # Linux
|
||||||
|
appdata = Path.home() / ".local" / "share"
|
||||||
|
|
||||||
|
providers_dir = appdata / "voicebox" / "providers"
|
||||||
|
providers_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
return providers_dir
|
||||||
|
|
||||||
|
|
||||||
|
def _get_provider_binary_name(provider_type: str) -> str:
|
||||||
|
"""Get the local binary filename for a provider type."""
|
||||||
|
system = platform.system()
|
||||||
|
ext = ".exe" if system == "Windows" else ""
|
||||||
|
|
||||||
|
binary_map = {
|
||||||
|
"pytorch-cpu": f"tts-provider-pytorch-cpu{ext}",
|
||||||
|
"pytorch-cuda": f"tts-provider-pytorch-cuda{ext}",
|
||||||
|
}
|
||||||
|
|
||||||
|
if provider_type not in binary_map:
|
||||||
|
raise ValueError(f"Unknown provider type: {provider_type}")
|
||||||
|
|
||||||
|
return binary_map[provider_type]
|
||||||
|
|
||||||
|
|
||||||
|
def _get_provider_download_name(provider_type: str) -> str:
|
||||||
|
"""Get the remote download filename for a provider type (includes platform suffix)."""
|
||||||
|
system = platform.system()
|
||||||
|
|
||||||
|
if system == "Windows":
|
||||||
|
platform_suffix = "windows"
|
||||||
|
ext = ".exe"
|
||||||
|
elif system == "Linux":
|
||||||
|
platform_suffix = "linux"
|
||||||
|
ext = ""
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Provider downloads not supported on {system}")
|
||||||
|
|
||||||
|
return f"tts-provider-{provider_type}-{platform_suffix}{ext}"
|
||||||
|
|
||||||
|
|
||||||
|
def _get_provider_download_url(provider_type: str) -> str:
|
||||||
|
"""Get the download URL for a provider."""
|
||||||
|
download_name = _get_provider_download_name(provider_type)
|
||||||
|
return f"{PROVIDER_DOWNLOAD_BASE_URL}/v{PROVIDER_VERSION}/{download_name}"
|
||||||
|
|
||||||
|
|
||||||
|
async def download_provider(provider_type: str) -> Path:
|
||||||
|
"""
|
||||||
|
Download a provider binary from Cloudflare R2.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
provider_type: Type of provider to download (e.g., "pytorch-cpu")
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Path to the downloaded provider binary
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If provider_type is invalid
|
||||||
|
httpx.HTTPError: If download fails
|
||||||
|
"""
|
||||||
|
if provider_type not in ["pytorch-cpu", "pytorch-cuda"]:
|
||||||
|
raise ValueError(f"Provider type {provider_type} cannot be downloaded")
|
||||||
|
|
||||||
|
progress_manager = get_progress_manager()
|
||||||
|
task_manager = get_task_manager()
|
||||||
|
|
||||||
|
binary_name = _get_provider_binary_name(provider_type)
|
||||||
|
download_url = _get_provider_download_url(provider_type)
|
||||||
|
destination = _get_providers_dir() / binary_name
|
||||||
|
|
||||||
|
# Start tracking download
|
||||||
|
task_manager.start_download(provider_type)
|
||||||
|
|
||||||
|
# Initialize progress state
|
||||||
|
progress_manager.update_progress(
|
||||||
|
model_name=provider_type,
|
||||||
|
current=0,
|
||||||
|
total=0, # Will be updated once we get Content-Length
|
||||||
|
filename=binary_name,
|
||||||
|
status="downloading",
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with httpx.AsyncClient(timeout=300.0) as client:
|
||||||
|
# First, get the file size
|
||||||
|
async with client.stream("GET", download_url) as response:
|
||||||
|
response.raise_for_status()
|
||||||
|
|
||||||
|
# Get total size from Content-Length header
|
||||||
|
total_size = int(response.headers.get("Content-Length", 0))
|
||||||
|
|
||||||
|
if total_size > 0:
|
||||||
|
progress_manager.update_progress(
|
||||||
|
model_name=provider_type,
|
||||||
|
current=0,
|
||||||
|
total=total_size,
|
||||||
|
filename=binary_name,
|
||||||
|
status="downloading",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Download with progress tracking
|
||||||
|
downloaded = 0
|
||||||
|
with open(destination, "wb") as f:
|
||||||
|
async for chunk in response.aiter_bytes(chunk_size=8192):
|
||||||
|
f.write(chunk)
|
||||||
|
downloaded += len(chunk)
|
||||||
|
|
||||||
|
# Update progress
|
||||||
|
progress_manager.update_progress(
|
||||||
|
model_name=provider_type,
|
||||||
|
current=downloaded,
|
||||||
|
total=total_size if total_size > 0 else downloaded,
|
||||||
|
filename=binary_name,
|
||||||
|
status="downloading",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Mark as complete
|
||||||
|
progress_manager.update_progress(
|
||||||
|
model_name=provider_type,
|
||||||
|
current=downloaded,
|
||||||
|
total=downloaded,
|
||||||
|
filename=binary_name,
|
||||||
|
status="complete",
|
||||||
|
)
|
||||||
|
task_manager.complete_download(provider_type)
|
||||||
|
|
||||||
|
# Make executable on Unix systems
|
||||||
|
if platform.system() != "Windows":
|
||||||
|
destination.chmod(0o755)
|
||||||
|
|
||||||
|
return destination
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
# Mark as error
|
||||||
|
progress_manager.update_progress(
|
||||||
|
model_name=provider_type,
|
||||||
|
current=0,
|
||||||
|
total=0,
|
||||||
|
filename=binary_name,
|
||||||
|
status="error",
|
||||||
|
)
|
||||||
|
task_manager.error_download(provider_type, str(e))
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
|
def get_provider_binary_path(provider_type: str) -> Optional[Path]:
|
||||||
|
"""
|
||||||
|
Get the path to an installed provider binary.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
provider_type: Type of provider
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Path to provider binary, or None if not installed
|
||||||
|
"""
|
||||||
|
binary_name = _get_provider_binary_name(provider_type)
|
||||||
|
provider_path = _get_providers_dir() / binary_name
|
||||||
|
|
||||||
|
if provider_path.exists() and provider_path.is_file():
|
||||||
|
return provider_path
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def delete_provider(provider_type: str) -> bool:
|
||||||
|
"""
|
||||||
|
Delete an installed provider binary.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
provider_type: Type of provider to delete
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
True if deleted, False if not found
|
||||||
|
"""
|
||||||
|
provider_path = get_provider_binary_path(provider_type)
|
||||||
|
|
||||||
|
if provider_path and provider_path.exists():
|
||||||
|
provider_path.unlink()
|
||||||
|
return True
|
||||||
|
|
||||||
|
return False
|
||||||
@@ -0,0 +1,187 @@
|
|||||||
|
"""
|
||||||
|
Local provider that communicates with standalone provider servers via HTTP.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Optional, Tuple
|
||||||
|
import base64
|
||||||
|
import io
|
||||||
|
import numpy as np
|
||||||
|
import httpx
|
||||||
|
import soundfile as sf
|
||||||
|
|
||||||
|
from .base import TTSProvider
|
||||||
|
from .types import ProviderHealth, ProviderStatus
|
||||||
|
|
||||||
|
|
||||||
|
class LocalProvider:
|
||||||
|
"""Provider that communicates with local subprocess via HTTP."""
|
||||||
|
|
||||||
|
def __init__(self, base_url: str):
|
||||||
|
"""
|
||||||
|
Initialize local provider.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
base_url: Base URL of the provider server (e.g., "http://localhost:8000")
|
||||||
|
"""
|
||||||
|
self.base_url = base_url.rstrip('/')
|
||||||
|
self.client = httpx.AsyncClient(timeout=300.0) # 5 minute timeout for generation
|
||||||
|
|
||||||
|
async def generate(
|
||||||
|
self,
|
||||||
|
text: str,
|
||||||
|
voice_prompt: dict,
|
||||||
|
language: str = "en",
|
||||||
|
seed: Optional[int] = None,
|
||||||
|
instruct: Optional[str] = None,
|
||||||
|
) -> Tuple[np.ndarray, int]:
|
||||||
|
"""Generate speech audio."""
|
||||||
|
response = await self.client.post(
|
||||||
|
f"{self.base_url}/tts/generate",
|
||||||
|
json={
|
||||||
|
"text": text,
|
||||||
|
"voice_prompt": voice_prompt,
|
||||||
|
"language": language,
|
||||||
|
"seed": seed,
|
||||||
|
"model_size": "1.7B", # TODO: Make configurable
|
||||||
|
}
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
data = response.json()
|
||||||
|
|
||||||
|
# Decode base64 audio
|
||||||
|
audio_bytes = base64.b64decode(data["audio"])
|
||||||
|
audio_buffer = io.BytesIO(audio_bytes)
|
||||||
|
audio, sample_rate = sf.read(audio_buffer)
|
||||||
|
|
||||||
|
return audio, data["sample_rate"]
|
||||||
|
|
||||||
|
async def create_voice_prompt(
|
||||||
|
self,
|
||||||
|
audio_path: str,
|
||||||
|
reference_text: str,
|
||||||
|
use_cache: bool = True,
|
||||||
|
) -> Tuple[dict, bool]:
|
||||||
|
"""Create voice prompt from reference audio."""
|
||||||
|
# Read audio file
|
||||||
|
with open(audio_path, 'rb') as f:
|
||||||
|
audio_data = f.read()
|
||||||
|
|
||||||
|
# Send multipart form data
|
||||||
|
files = {
|
||||||
|
"audio": ("audio.wav", audio_data, "audio/wav")
|
||||||
|
}
|
||||||
|
data = {
|
||||||
|
"reference_text": reference_text,
|
||||||
|
"use_cache": str(use_cache).lower(),
|
||||||
|
}
|
||||||
|
|
||||||
|
response = await self.client.post(
|
||||||
|
f"{self.base_url}/tts/create_voice_prompt",
|
||||||
|
files=files,
|
||||||
|
data=data,
|
||||||
|
)
|
||||||
|
response.raise_for_status()
|
||||||
|
result = response.json()
|
||||||
|
|
||||||
|
return result["voice_prompt"], result.get("was_cached", False)
|
||||||
|
|
||||||
|
async def combine_voice_prompts(
|
||||||
|
self,
|
||||||
|
audio_paths: list[str],
|
||||||
|
reference_texts: list[str],
|
||||||
|
) -> Tuple[np.ndarray, str]:
|
||||||
|
"""
|
||||||
|
Combine multiple voice prompts.
|
||||||
|
|
||||||
|
Note: This is not implemented in the provider API yet.
|
||||||
|
For now, we'll combine locally by concatenating audio.
|
||||||
|
"""
|
||||||
|
import numpy as np
|
||||||
|
from ..utils.audio import load_audio, normalize_audio
|
||||||
|
|
||||||
|
combined_audio = []
|
||||||
|
for audio_path in audio_paths:
|
||||||
|
audio, sr = load_audio(audio_path)
|
||||||
|
audio = normalize_audio(audio)
|
||||||
|
combined_audio.append(audio)
|
||||||
|
|
||||||
|
# Concatenate audio
|
||||||
|
mixed = np.concatenate(combined_audio)
|
||||||
|
mixed = normalize_audio(mixed)
|
||||||
|
|
||||||
|
# Combine texts
|
||||||
|
combined_text = " ".join(reference_texts)
|
||||||
|
|
||||||
|
return mixed, combined_text
|
||||||
|
|
||||||
|
async def load_model(self, model_size: str) -> None:
|
||||||
|
"""Load TTS model."""
|
||||||
|
# Model loading is handled automatically by the provider server
|
||||||
|
# when generate() is called, so this is a no-op
|
||||||
|
pass
|
||||||
|
|
||||||
|
def unload_model(self) -> None:
|
||||||
|
"""Unload model to free memory."""
|
||||||
|
# Model unloading is handled by the provider server
|
||||||
|
# This is a no-op for local providers
|
||||||
|
pass
|
||||||
|
|
||||||
|
def is_loaded(self) -> bool:
|
||||||
|
"""Check if model is loaded."""
|
||||||
|
# We can't know this without querying the provider
|
||||||
|
# Return True optimistically
|
||||||
|
return True
|
||||||
|
|
||||||
|
def _get_model_path(self, model_size: str) -> str:
|
||||||
|
"""Get model path for a given size."""
|
||||||
|
# For local providers, model paths are handled by the provider server
|
||||||
|
# Return a placeholder
|
||||||
|
return f"Qwen/Qwen3-TTS-12Hz-{model_size}-Base"
|
||||||
|
|
||||||
|
async def health(self) -> ProviderHealth:
|
||||||
|
"""Get provider health status."""
|
||||||
|
try:
|
||||||
|
response = await self.client.get(f"{self.base_url}/tts/health")
|
||||||
|
response.raise_for_status()
|
||||||
|
data = response.json()
|
||||||
|
return ProviderHealth(
|
||||||
|
status=data["status"],
|
||||||
|
provider=data["provider"],
|
||||||
|
version=data.get("version"),
|
||||||
|
model=data.get("model"),
|
||||||
|
device=data.get("device"),
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
return ProviderHealth(
|
||||||
|
status="unhealthy",
|
||||||
|
provider="local",
|
||||||
|
version=None,
|
||||||
|
model=None,
|
||||||
|
device=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def status(self) -> ProviderStatus:
|
||||||
|
"""Get provider model status."""
|
||||||
|
try:
|
||||||
|
response = await self.client.get(f"{self.base_url}/tts/status")
|
||||||
|
response.raise_for_status()
|
||||||
|
data = response.json()
|
||||||
|
return ProviderStatus(
|
||||||
|
model_loaded=data["model_loaded"],
|
||||||
|
model_size=data.get("model_size"),
|
||||||
|
available_sizes=data.get("available_sizes", []),
|
||||||
|
gpu_available=data.get("gpu_available"),
|
||||||
|
vram_used_mb=data.get("vram_used_mb"),
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
return ProviderStatus(
|
||||||
|
model_loaded=False,
|
||||||
|
model_size=None,
|
||||||
|
available_sizes=[],
|
||||||
|
gpu_available=None,
|
||||||
|
vram_used_mb=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def stop(self) -> None:
|
||||||
|
"""Stop the provider (close HTTP client)."""
|
||||||
|
await self.client.aclose()
|
||||||
@@ -0,0 +1,34 @@
|
|||||||
|
"""
|
||||||
|
Shared types for TTS providers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Optional, TypedDict
|
||||||
|
from enum import Enum
|
||||||
|
|
||||||
|
|
||||||
|
class ProviderType(str, Enum):
|
||||||
|
"""Available provider types."""
|
||||||
|
BUNDLED_MLX = "bundled-mlx"
|
||||||
|
BUNDLED_PYTORCH = "bundled-pytorch"
|
||||||
|
PYTORCH_CPU = "pytorch-cpu"
|
||||||
|
PYTORCH_CUDA = "pytorch-cuda"
|
||||||
|
REMOTE = "remote"
|
||||||
|
OPENAI = "openai"
|
||||||
|
|
||||||
|
|
||||||
|
class ProviderHealth(TypedDict):
|
||||||
|
"""Provider health status."""
|
||||||
|
status: str # "healthy", "unhealthy", "starting"
|
||||||
|
provider: str
|
||||||
|
version: Optional[str]
|
||||||
|
model: Optional[str]
|
||||||
|
device: Optional[str]
|
||||||
|
|
||||||
|
|
||||||
|
class ProviderStatus(TypedDict):
|
||||||
|
"""Provider model status."""
|
||||||
|
model_loaded: bool
|
||||||
|
model_size: Optional[str]
|
||||||
|
available_sizes: list[str]
|
||||||
|
gpu_available: Optional[bool]
|
||||||
|
vram_used_mb: Optional[int]
|
||||||
+36
-16
@@ -1,5 +1,5 @@
|
|||||||
"""
|
"""
|
||||||
TTS inference module - delegates to backend abstraction layer.
|
TTS inference module - delegates to provider abstraction layer.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
@@ -7,31 +7,51 @@ import numpy as np
|
|||||||
import io
|
import io
|
||||||
import soundfile as sf
|
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:
|
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():
|
def unload_tts_model():
|
||||||
"""Unload TTS model to free memory."""
|
"""Unload TTS model to free memory."""
|
||||||
backend = get_tts_backend()
|
manager = get_provider_manager()
|
||||||
backend.unload_model()
|
provider = manager._get_default_provider()
|
||||||
|
provider.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()
|
|
||||||
|
|
||||||
|
|
||||||
def audio_to_wav_bytes(audio: np.ndarray, sample_rate: int) -> bytes:
|
def audio_to_wav_bytes(audio: np.ndarray, sample_rate: int) -> bytes:
|
||||||
|
|||||||
@@ -10,14 +10,17 @@
|
|||||||
|
|
||||||
Split the monolithic backend into modular components:
|
Split the monolithic backend into modular components:
|
||||||
|
|
||||||
1. **Main App** (~150-200MB): Tauri + FastAPI backend + Whisper + UI/profiles/history
|
1. **Main App**:
|
||||||
2. **TTS Providers** (downloadable plugins): Separate executables for model inference
|
- 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:
|
This architecture solves:
|
||||||
|
|
||||||
- ✅ GitHub 2GB release artifact limit
|
- ✅ GitHub 2GB release artifact limit
|
||||||
- ✅ Frequent app updates without re-downloading large python binaries
|
- ✅ Frequent app updates without re-downloading large python binaries (Windows/Linux)
|
||||||
- ✅ User choice of compute backend (CPU/GPU/Cloud)
|
- ✅ 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)
|
- ✅ External provider support (OpenAI, custom servers)
|
||||||
- ✅ Future extensibility
|
- ✅ Future extensibility
|
||||||
|
|
||||||
@@ -25,6 +28,7 @@ This architecture solves:
|
|||||||
|
|
||||||
## Architecture Diagram
|
## Architecture Diagram
|
||||||
|
|
||||||
|
### Windows / Linux
|
||||||
```
|
```
|
||||||
┌─────────────────────────────────────────────────────────┐
|
┌─────────────────────────────────────────────────────────┐
|
||||||
│ Voicebox App (Tauri + Backend) ~150MB │
|
│ Voicebox App (Tauri + Backend) ~150MB │
|
||||||
@@ -39,27 +43,43 @@ This architecture solves:
|
|||||||
│
|
│
|
||||||
HTTP/IPC │
|
HTTP/IPC │
|
||||||
│
|
│
|
||||||
┌────────────────────────────────┼─────────────────┐
|
┌─────────────────────┴─────────────────┐
|
||||||
│ │ │
|
│ │
|
||||||
▼ ▼ ▼
|
▼ ▼
|
||||||
┌─────────────────┐ ┌─────────────────┐ ┌──────────────────┐
|
┌─────────────────────┐ ┌─────────────────────┐
|
||||||
│ TTS Provider: │ │ TTS Provider: │ │ TTS Provider: │
|
│ TTS Provider: │ │ TTS Provider: │
|
||||||
│ PyTorch CPU │ │ PyTorch CUDA │ │ MLX (Apple) │
|
│ PyTorch CPU │ │ PyTorch CUDA │
|
||||||
│ │ │ │ │ │
|
│ │ │ │
|
||||||
│ ~300MB │ │ ~2.4GB │ │ ~800MB │
|
│ ~300MB │ │ ~2.4GB │
|
||||||
│ │ │ │ │ │
|
│ │ │ │
|
||||||
│ Local inference │ │ GPU inference │ │ Metal inference │
|
│ Local inference │ │ GPU inference │
|
||||||
└─────────────────┘ └─────────────────┘ └──────────────────┘
|
└─────────────────────┘ └─────────────────────┘
|
||||||
│ │ │
|
│ │
|
||||||
└────────────────────────┴─────────────────────┘
|
└───────────────┬───────────────────────┘
|
||||||
│
|
│
|
||||||
┌─────────────▼──────────────┐
|
┌─────────────▼──────────────┐
|
||||||
│ Future Providers: │
|
│ Future Providers: │
|
||||||
│ • Remote Server │
|
│ • Remote Server │
|
||||||
│ • OpenAI API │
|
│ • OpenAI API │
|
||||||
│ • ElevenLabs │
|
│ • ElevenLabs │
|
||||||
│ • Custom Docker Container │
|
│ • 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)
|
#### 1. Main App (voicebox.exe / .app / .AppImage)
|
||||||
|
|
||||||
**Size:** ~100-150MB
|
**Windows/Linux Size:** ~100-150MB
|
||||||
|
**macOS Size:** ~300-350MB (includes MLX)
|
||||||
|
|
||||||
**Includes:**
|
**Includes:**
|
||||||
|
|
||||||
- Tauri runtime + React UI
|
- Tauri runtime + React UI
|
||||||
- FastAPI backend (pure Python, no PyTorch)
|
- FastAPI backend (pure Python, no PyTorch on Windows/Linux)
|
||||||
- Whisper model (tiny, ~50MB)
|
- Whisper model (tiny, ~50MB)
|
||||||
- SQLite database
|
- SQLite database
|
||||||
- Profile/history/audio editing logic
|
- 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)
|
- PyTorch (CPU or CUDA)
|
||||||
- TTS models (Qwen3-TTS)
|
- TTS models (Qwen3-TTS)
|
||||||
@@ -147,23 +169,7 @@ This architecture solves:
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
#### 4. TTS Provider: MLX
|
#### 4. TTS Provider: Remote
|
||||||
|
|
||||||
**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
|
|
||||||
|
|
||||||
**Binary:** None (built-in config)
|
**Binary:** None (built-in config)
|
||||||
**Size:** 0MB
|
**Size:** 0MB
|
||||||
@@ -182,7 +188,7 @@ This architecture solves:
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
#### 6. TTS Provider: OpenAI
|
#### 5. TTS Provider: OpenAI
|
||||||
|
|
||||||
**Binary:** None (API wrapper)
|
**Binary:** None (API wrapper)
|
||||||
**Size:** 0MB
|
**Size:** 0MB
|
||||||
@@ -296,7 +302,10 @@ Model status.
|
|||||||
|
|
||||||
```python
|
```python
|
||||||
class ProviderManager:
|
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):
|
def __init__(self):
|
||||||
self.active_provider: Optional[Provider] = None
|
self.active_provider: Optional[Provider] = None
|
||||||
@@ -308,8 +317,6 @@ class ProviderManager:
|
|||||||
return await self._start_local_provider("tts-provider-pytorch-cpu.exe")
|
return await self._start_local_provider("tts-provider-pytorch-cpu.exe")
|
||||||
elif provider_type == "pytorch-cuda":
|
elif provider_type == "pytorch-cuda":
|
||||||
return await self._start_local_provider("tts-provider-pytorch-cuda.exe")
|
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":
|
elif provider_type == "remote":
|
||||||
return self.config["remote_url"]
|
return self.config["remote_url"]
|
||||||
elif provider_type == "openai":
|
elif provider_type == "openai":
|
||||||
@@ -434,15 +441,14 @@ class OpenAIProvider(TTSProvider):
|
|||||||
|
|
||||||
```python
|
```python
|
||||||
class ProviderInstaller:
|
class ProviderInstaller:
|
||||||
"""Handles provider download and installation."""
|
"""Handles provider download and installation (Windows/Linux only)."""
|
||||||
|
|
||||||
async def download_provider(self, provider_type: str):
|
async def download_provider(self, provider_type: str):
|
||||||
"""Download provider binary from R2."""
|
"""Download provider binary from R2."""
|
||||||
|
|
||||||
binary_name = {
|
binary_name = {
|
||||||
"pytorch-cpu": "tts-provider-pytorch-cpu.exe",
|
"pytorch-cpu": "tts-provider-pytorch-cpu.exe",
|
||||||
"pytorch-cuda": "tts-provider-pytorch-cuda.exe",
|
"pytorch-cuda": "tts-provider-pytorch-cuda.exe"
|
||||||
"mlx": "tts-provider-mlx"
|
|
||||||
}[provider_type]
|
}[provider_type]
|
||||||
|
|
||||||
download_url = f"https://downloads.voicebox.sh/providers/v{PROVIDER_VERSION}/{binary_name}"
|
download_url = f"https://downloads.voicebox.sh/providers/v{PROVIDER_VERSION}/{binary_name}"
|
||||||
@@ -525,44 +531,38 @@ export function ProviderSettings() {
|
|||||||
)}
|
)}
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
{/* PyTorch CPU */}
|
{/* PyTorch CPU (Windows/Linux only) */}
|
||||||
<div className="flex items-center justify-between">
|
{!isMacOS && (
|
||||||
<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 && (
|
|
||||||
<div className="flex items-center justify-between">
|
<div className="flex items-center justify-between">
|
||||||
<div className="flex items-center space-x-2">
|
<div className="flex items-center space-x-2">
|
||||||
<RadioGroupItem value="mlx" id="mlx" />
|
<RadioGroupItem value="pytorch-cpu" id="cpu" />
|
||||||
<Label htmlFor="mlx">
|
<Label htmlFor="cpu">
|
||||||
<div className="font-medium">MLX (Apple Silicon)</div>
|
<div className="font-medium">PyTorch CPU</div>
|
||||||
<div className="text-sm text-muted-foreground">
|
<div className="text-sm text-muted-foreground">
|
||||||
Optimized for M1/M2/M3 chips
|
Works on any system, slower inference
|
||||||
</div>
|
</div>
|
||||||
</Label>
|
</Label>
|
||||||
</div>
|
</div>
|
||||||
{!installedProviders?.includes("mlx") && (
|
{!installedProviders?.includes("pytorch-cpu") && (
|
||||||
<Button onClick={() => downloadProvider("mlx")} size="sm">
|
<Button onClick={() => downloadProvider("pytorch-cpu")} size="sm">
|
||||||
Download (800MB)
|
Download (300MB)
|
||||||
</Button>
|
</Button>
|
||||||
)}
|
)}
|
||||||
</div>
|
</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 */}
|
{/* Remote */}
|
||||||
<div className="space-y-2">
|
<div className="space-y-2">
|
||||||
<div className="flex items-center space-x-2">
|
<div className="flex items-center space-x-2">
|
||||||
@@ -608,14 +608,18 @@ export function ProviderSettings() {
|
|||||||
```
|
```
|
||||||
voicebox/
|
voicebox/
|
||||||
├── backend/
|
├── 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/
|
│ ├── providers/
|
||||||
│ │ ├── __init__.py # ProviderManager
|
│ │ ├── __init__.py # ProviderManager (Windows/Linux)
|
||||||
│ │ ├── base.py # TTSProvider ABC
|
│ │ ├── base.py # TTSProvider Protocol
|
||||||
│ │ ├── local.py # LocalProvider (subprocess)
|
│ │ ├── local.py # LocalProvider (subprocess)
|
||||||
│ │ ├── remote.py # RemoteProvider (HTTP)
|
│ │ ├── remote.py # RemoteProvider (HTTP)
|
||||||
│ │ ├── openai.py # OpenAIProvider (API wrapper)
|
│ │ ├── openai.py # OpenAIProvider (API wrapper)
|
||||||
│ │ └── installer.py # Provider download logic
|
│ │ └── installer.py # Provider download logic (Windows/Linux)
|
||||||
│ ├── profiles.py # Voice profile management
|
│ ├── profiles.py # Voice profile management
|
||||||
│ ├── history.py # Generation history
|
│ ├── history.py # Generation history
|
||||||
│ ├── transcribe.py # Whisper (still bundled)
|
│ ├── transcribe.py # Whisper (still bundled)
|
||||||
@@ -628,27 +632,22 @@ voicebox/
|
|||||||
│ │ ├── requirements.txt # torch (CPU), qwen-tts, transformers
|
│ │ ├── requirements.txt # torch (CPU), qwen-tts, transformers
|
||||||
│ │ └── build.spec # PyInstaller spec
|
│ │ └── build.spec # PyInstaller spec
|
||||||
│ │
|
│ │
|
||||||
│ ├── pytorch-cuda/
|
│ └── 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/
|
|
||||||
│ ├── main.py # FastAPI server for TTS
|
│ ├── main.py # FastAPI server for TTS
|
||||||
│ ├── mlx_backend.py # MLX TTS logic
|
│ ├── tts_backend.py # PyTorch TTS logic
|
||||||
│ ├── requirements.txt # mlx, qwen-tts-mlx
|
│ ├── requirements.txt # torch+cu121, qwen-tts, transformers
|
||||||
│ └── build.spec # PyInstaller spec
|
│ └── build.spec # PyInstaller spec
|
||||||
│
|
│
|
||||||
├── app/ # Frontend (Tauri + React)
|
├── app/ # Frontend (Tauri + React)
|
||||||
│ └── src/
|
│ └── src/
|
||||||
│ └── components/
|
│ └── components/
|
||||||
│ └── ServerSettings/
|
│ └── ServerSettings/
|
||||||
│ └── ProviderSettings.tsx
|
│ └── ProviderSettings.tsx # Only shown on Windows/Linux
|
||||||
│
|
│
|
||||||
└── tauri/
|
└── tauri/
|
||||||
└── src-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
|
### 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
|
1. Create separate PyInstaller specs for each provider
|
||||||
2. Build provider executables:
|
2. Build provider executables:
|
||||||
- `tts-provider-pytorch-cpu.exe` (~300MB)
|
- `tts-provider-pytorch-cpu.exe` (~300MB)
|
||||||
- `tts-provider-pytorch-cuda.exe` (~2.4GB)
|
- `tts-provider-pytorch-cuda.exe` (~2.4GB)
|
||||||
- `tts-provider-mlx` (~800MB, macOS)
|
|
||||||
3. Test subprocess communication
|
3. Test subprocess communication
|
||||||
4. Upload providers to Cloudflare R2
|
4. Upload providers to Cloudflare R2
|
||||||
|
|
||||||
**Result:** Provider binaries exist but aren't used yet
|
**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
|
### 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
|
1. Exclude PyTorch/Qwen3-TTS from Windows/Linux main app PyInstaller spec
|
||||||
2. Main app now requires provider download
|
2. Windows/Linux app now requires provider download
|
||||||
3. Update GitHub CI to build multiple artifacts:
|
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-cpu-{version}.exe`
|
||||||
- `tts-provider-pytorch-cuda-{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
|
## User Flows
|
||||||
|
|
||||||
### First-Time Setup
|
### First-Time Setup (Windows/Linux)
|
||||||
|
|
||||||
1. User downloads and installs Voicebox (~150MB)
|
1. User downloads and installs Voicebox (~150MB)
|
||||||
2. App launches → detects no TTS provider installed
|
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
|
✓ Works on any system
|
||||||
✗ Slower inference
|
✗ Slower inference
|
||||||
|
|
||||||
[ ] MLX (800MB) [Download]
|
|
||||||
✓ Fast on Apple Silicon
|
|
||||||
✗ macOS only (M1/M2/M3)
|
|
||||||
|
|
||||||
[ ] Remote Server
|
[ ] Remote Server
|
||||||
URL: ___________________
|
URL: ___________________
|
||||||
|
|
||||||
@@ -799,19 +796,31 @@ async def check_provider_compatibility(provider_version: str) -> bool:
|
|||||||
5. Provider installs to AppData/Application Support
|
5. Provider installs to AppData/Application Support
|
||||||
6. App starts provider → ready to use
|
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)
|
### App Update Flow (No Provider Change)
|
||||||
|
|
||||||
**Scenario:** Bug fix in UI, no backend changes
|
**Scenario:** Bug fix in UI, no backend changes
|
||||||
|
|
||||||
|
**Windows/Linux:**
|
||||||
1. User gets update notification: "Voicebox v0.2.1 available"
|
1. User gets update notification: "Voicebox v0.2.1 available"
|
||||||
2. Downloads update (~150MB, not 2.4GB!)
|
2. Downloads update (~150MB, not 2.4GB!)
|
||||||
3. Installs and restarts
|
3. Installs and restarts
|
||||||
4. **Provider stays the same** (no re-download needed)
|
4. **Provider stays the same** (no re-download needed)
|
||||||
5. App starts using existing provider
|
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 |
|
| 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 |
|
| **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 |
|
| **External Provider Support** | Users can run their own TTS servers |
|
||||||
| **Bandwidth Savings** | Only download provider once, app updates are small |
|
| **Bandwidth Savings** | Only download provider once, app updates are small |
|
||||||
| **Future-Proof** | Easy to add new providers (ElevenLabs, custom models) |
|
| **Future-Proof** | Easy to add new providers (ElevenLabs, custom models) |
|
||||||
|
|||||||
@@ -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.
|
||||||
@@ -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()
|
||||||
@@ -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()
|
||||||
@@ -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
|
||||||
@@ -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()
|
||||||
@@ -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()
|
||||||
@@ -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
|
||||||
Reference in New Issue
Block a user