mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-09-15 12:50:42 -07:00
- Added support for MLX backend on Apple Silicon, enabling optimized performance for TTS and STT tasks. - Implemented platform detection to dynamically select between MLX and PyTorch based on the runtime environment. - Updated build process to include MLX-specific dependencies and configurations for macOS. - Refactored backend code to improve model loading and inference logic, accommodating backend-specific requirements. - Enhanced documentation to clarify backend selection and performance benefits for different platforms. - Streamlined installation instructions and troubleshooting guidance for MLX-related issues.
167 lines
3.9 KiB
Python
167 lines
3.9 KiB
Python
"""
|
|
Backend abstraction layer for TTS and STT.
|
|
|
|
Provides a unified interface for MLX and PyTorch backends.
|
|
"""
|
|
|
|
from typing import Protocol, Optional, Tuple, List
|
|
from typing_extensions import runtime_checkable
|
|
import numpy as np
|
|
|
|
from ..platform_detect import get_backend_type
|
|
|
|
|
|
@runtime_checkable
|
|
class TTSBackend(Protocol):
|
|
"""Protocol for TTS backend implementations."""
|
|
|
|
async def load_model(self, model_size: str) -> None:
|
|
"""Load TTS model."""
|
|
...
|
|
|
|
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.
|
|
|
|
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.
|
|
|
|
Returns:
|
|
Tuple of (combined_audio_array, combined_text)
|
|
"""
|
|
...
|
|
|
|
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 audio from text.
|
|
|
|
Returns:
|
|
Tuple of (audio_array, sample_rate)
|
|
"""
|
|
...
|
|
|
|
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.
|
|
|
|
Returns:
|
|
Model path or HuggingFace Hub ID
|
|
"""
|
|
...
|
|
|
|
|
|
@runtime_checkable
|
|
class STTBackend(Protocol):
|
|
"""Protocol for STT (Speech-to-Text) backend implementations."""
|
|
|
|
async def load_model(self, model_size: str) -> None:
|
|
"""Load STT model."""
|
|
...
|
|
|
|
async def transcribe(
|
|
self,
|
|
audio_path: str,
|
|
language: Optional[str] = None,
|
|
) -> str:
|
|
"""
|
|
Transcribe audio to text.
|
|
|
|
Returns:
|
|
Transcribed text
|
|
"""
|
|
...
|
|
|
|
def unload_model(self) -> None:
|
|
"""Unload model to free memory."""
|
|
...
|
|
|
|
def is_loaded(self) -> bool:
|
|
"""Check if model is loaded."""
|
|
...
|
|
|
|
|
|
# Global backend instances
|
|
_tts_backend: Optional[TTSBackend] = None
|
|
_stt_backend: Optional[STTBackend] = None
|
|
|
|
|
|
def get_tts_backend() -> TTSBackend:
|
|
"""
|
|
Get or create TTS backend instance based on platform.
|
|
|
|
Returns:
|
|
TTS backend instance (MLX or PyTorch)
|
|
"""
|
|
global _tts_backend
|
|
|
|
if _tts_backend is None:
|
|
backend_type = get_backend_type()
|
|
|
|
if backend_type == "mlx":
|
|
from .mlx_backend import MLXTTSBackend
|
|
_tts_backend = MLXTTSBackend()
|
|
else:
|
|
from .pytorch_backend import PyTorchTTSBackend
|
|
_tts_backend = PyTorchTTSBackend()
|
|
|
|
return _tts_backend
|
|
|
|
|
|
def get_stt_backend() -> STTBackend:
|
|
"""
|
|
Get or create STT backend instance based on platform.
|
|
|
|
Returns:
|
|
STT backend instance (MLX or PyTorch)
|
|
"""
|
|
global _stt_backend
|
|
|
|
if _stt_backend is None:
|
|
backend_type = get_backend_type()
|
|
|
|
if backend_type == "mlx":
|
|
from .mlx_backend import MLXSTTBackend
|
|
_stt_backend = MLXSTTBackend()
|
|
else:
|
|
from .pytorch_backend import PyTorchSTTBackend
|
|
_stt_backend = PyTorchSTTBackend()
|
|
|
|
return _stt_backend
|
|
|
|
|
|
def reset_backends():
|
|
"""Reset backend instances (useful for testing)."""
|
|
global _tts_backend, _stt_backend
|
|
_tts_backend = None
|
|
_stt_backend = None
|