mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-09-26 13:45:16 -07:00
refactor: remove dead code, deduplicate backends
Phase 1 - delete dead code: - studio.py, migrate_add_instruct.py, utils/validation.py - duplicate _profile_to_response in main.py, duplicate asyncio import - pointless _get_profiles_dir/_get_generations_dir wrappers - duplicate LANGUAGE_CODE_TO_NAME and WHISPER_HF_REPOS constants Phase 2 - extract backends/base.py with shared utilities: - is_model_cached() replaces 7 copy-pasted HF cache checks - get_torch_device() replaces 5 device detection methods - combine_voice_prompts() replaces 5 identical implementations - model_load_progress() ctx manager replaces progress boilerplate in all backends - patch_chatterbox_f32() replaces identical monkey-patches in both chatterbox backends net -1078 lines across the backend
This commit is contained in:
+31
-20
@@ -17,23 +17,38 @@ Production-quality FastAPI backend for Qwen3-TTS voice cloning.
|
|||||||
|
|
||||||
```
|
```
|
||||||
backend/
|
backend/
|
||||||
├── main.py # FastAPI app with all routes
|
├── main.py # FastAPI app with all routes
|
||||||
├── models.py # Pydantic request/response models
|
├── server.py # PyInstaller entry point, CLI arg parsing
|
||||||
├── platform_detect.py # Platform detection for backend selection
|
├── models.py # Pydantic request/response models
|
||||||
├── tts.py # TTS backend abstraction (delegates to MLX or PyTorch)
|
├── config.py # Data directory configuration
|
||||||
├── transcribe.py # STT backend abstraction (delegates to MLX or PyTorch)
|
├── database.py # SQLAlchemy ORM models + migrations
|
||||||
├── backends/ # Backend implementations
|
├── platform_detect.py # Platform detection for backend selection
|
||||||
│ ├── __init__.py # Backend factory and protocols
|
├── tts.py # TTS backend facade
|
||||||
│ ├── mlx_backend.py # MLX backend (Apple Silicon)
|
├── transcribe.py # STT backend facade
|
||||||
│ └── pytorch_backend.py # PyTorch backend (Windows/Linux/Intel)
|
├── profiles.py # Voice profile CRUD
|
||||||
├── profiles.py # Voice profile CRUD
|
├── history.py # Generation history CRUD
|
||||||
├── history.py # Generation history
|
├── channels.py # Audio channel management
|
||||||
├── studio.py # Audio editing (TODO)
|
├── stories.py # Story/timeline management + audio export
|
||||||
├── database.py # SQLite ORM
|
├── effects.py # Effect preset CRUD
|
||||||
|
├── versions.py # Generation version management
|
||||||
|
├── export_import.py # ZIP export/import for profiles and generations
|
||||||
|
├── backends/ # Backend implementations
|
||||||
|
│ ├── __init__.py # Protocols, model config registry, factory functions
|
||||||
|
│ ├── mlx_backend.py # MLX backend (Apple Silicon)
|
||||||
|
│ ├── pytorch_backend.py # PyTorch backend (Windows/Linux/Intel)
|
||||||
|
│ ├── chatterbox_backend.py # Chatterbox Multilingual TTS
|
||||||
|
│ ├── chatterbox_turbo_backend.py # Chatterbox Turbo TTS
|
||||||
|
│ └── luxtts_backend.py # LuxTTS backend
|
||||||
└── utils/
|
└── utils/
|
||||||
├── audio.py # Audio processing utilities
|
├── audio.py # Audio load/save/normalize/validate/trim
|
||||||
├── cache.py # Voice prompt caching
|
├── cache.py # Voice prompt caching (memory + disk)
|
||||||
└── validation.py # Input validation
|
├── effects.py # Audio effects engine (pedalboard)
|
||||||
|
├── chunked_tts.py # Text chunking + audio concatenation
|
||||||
|
├── progress.py # SSE progress tracking
|
||||||
|
├── tasks.py # Active task tracking
|
||||||
|
├── hf_progress.py # HuggingFace download progress tracking
|
||||||
|
├── hf_offline_patch.py # HuggingFace offline mode patch (MLX)
|
||||||
|
└── images.py # Avatar image processing
|
||||||
```
|
```
|
||||||
|
|
||||||
### Backend Selection
|
### Backend Selection
|
||||||
@@ -450,12 +465,8 @@ Error responses include details:
|
|||||||
|
|
||||||
- [ ] WebSocket support for generation progress
|
- [ ] WebSocket support for generation progress
|
||||||
- [ ] Batch generation endpoint
|
- [ ] Batch generation endpoint
|
||||||
- [ ] Audio effects (M3GAN, etc.)
|
|
||||||
- [ ] Voice design (text-to-voice)
|
- [ ] Voice design (text-to-voice)
|
||||||
- [ ] Audio studio timeline features
|
|
||||||
- [ ] Project management
|
|
||||||
- [ ] Authentication & rate limiting
|
- [ ] Authentication & rate limiting
|
||||||
- [ ] Export/import profiles
|
|
||||||
|
|
||||||
## License
|
## License
|
||||||
|
|
||||||
|
|||||||
@@ -13,6 +13,20 @@ import numpy as np
|
|||||||
|
|
||||||
from ..platform_detect import get_backend_type
|
from ..platform_detect import get_backend_type
|
||||||
|
|
||||||
|
LANGUAGE_CODE_TO_NAME = {
|
||||||
|
"zh": "chinese", "en": "english", "ja": "japanese", "ko": "korean",
|
||||||
|
"de": "german", "fr": "french", "ru": "russian", "pt": "portuguese",
|
||||||
|
"es": "spanish", "it": "italian",
|
||||||
|
}
|
||||||
|
|
||||||
|
WHISPER_HF_REPOS = {
|
||||||
|
"base": "openai/whisper-base",
|
||||||
|
"small": "openai/whisper-small",
|
||||||
|
"medium": "openai/whisper-medium",
|
||||||
|
"large": "openai/whisper-large-v3",
|
||||||
|
"turbo": "openai/whisper-large-v3-turbo",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ModelConfig:
|
class ModelConfig:
|
||||||
|
|||||||
@@ -0,0 +1,277 @@
|
|||||||
|
"""
|
||||||
|
Shared utilities for TTS/STT backend implementations.
|
||||||
|
|
||||||
|
Eliminates duplication of cache checking, device detection,
|
||||||
|
voice prompt combination, and model loading progress tracking.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import platform
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Callable, List, Optional, Tuple
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
from ..utils.audio import normalize_audio, load_audio
|
||||||
|
from ..utils.progress import get_progress_manager
|
||||||
|
from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback
|
||||||
|
from ..utils.tasks import get_task_manager
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# HuggingFace cache checking
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def is_model_cached(
|
||||||
|
hf_repo: str,
|
||||||
|
*,
|
||||||
|
weight_extensions: tuple[str, ...] = (".safetensors", ".bin"),
|
||||||
|
required_files: Optional[list[str]] = None,
|
||||||
|
) -> bool:
|
||||||
|
"""
|
||||||
|
Check if a HuggingFace model is fully cached locally.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
hf_repo: HuggingFace repo ID (e.g. "Qwen/Qwen3-TTS-12Hz-1.7B-Base")
|
||||||
|
weight_extensions: File extensions that count as model weights.
|
||||||
|
required_files: If set, check that these specific filenames exist
|
||||||
|
in snapshots instead of checking by extension.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
True if model is fully cached, False if missing or incomplete.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
from huggingface_hub import constants as hf_constants
|
||||||
|
|
||||||
|
repo_cache = Path(hf_constants.HF_HUB_CACHE) / (
|
||||||
|
"models--" + hf_repo.replace("/", "--")
|
||||||
|
)
|
||||||
|
|
||||||
|
if not repo_cache.exists():
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Incomplete blobs mean a download is still in progress
|
||||||
|
blobs_dir = repo_cache / "blobs"
|
||||||
|
if blobs_dir.exists() and any(blobs_dir.glob("*.incomplete")):
|
||||||
|
logger.debug(f"Found .incomplete files for {hf_repo}")
|
||||||
|
return False
|
||||||
|
|
||||||
|
snapshots_dir = repo_cache / "snapshots"
|
||||||
|
if not snapshots_dir.exists():
|
||||||
|
return False
|
||||||
|
|
||||||
|
if required_files:
|
||||||
|
# Check that every required filename exists somewhere in snapshots
|
||||||
|
for fname in required_files:
|
||||||
|
if not any(snapshots_dir.rglob(fname)):
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
# Check that at least one weight file exists
|
||||||
|
for ext in weight_extensions:
|
||||||
|
if any(snapshots_dir.rglob(f"*{ext}")):
|
||||||
|
return True
|
||||||
|
|
||||||
|
logger.debug(f"No model weights found for {hf_repo}")
|
||||||
|
return False
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"Error checking cache for {hf_repo}: {e}")
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Device detection
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def get_torch_device(
|
||||||
|
*,
|
||||||
|
allow_xpu: bool = False,
|
||||||
|
allow_directml: bool = False,
|
||||||
|
allow_mps: bool = False,
|
||||||
|
force_cpu_on_mac: bool = False,
|
||||||
|
) -> str:
|
||||||
|
"""
|
||||||
|
Detect the best available torch device.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
allow_xpu: Check for Intel XPU (IPEX) support.
|
||||||
|
allow_directml: Check for DirectML (Windows) support.
|
||||||
|
allow_mps: Allow MPS (Apple Silicon). If False, MPS falls back to CPU.
|
||||||
|
force_cpu_on_mac: Force CPU on macOS regardless of GPU availability.
|
||||||
|
"""
|
||||||
|
if force_cpu_on_mac and platform.system() == "Darwin":
|
||||||
|
return "cpu"
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
if torch.cuda.is_available():
|
||||||
|
return "cuda"
|
||||||
|
|
||||||
|
if allow_xpu:
|
||||||
|
try:
|
||||||
|
import intel_extension_for_pytorch # noqa: F401
|
||||||
|
if hasattr(torch, "xpu") and torch.xpu.is_available():
|
||||||
|
return "xpu"
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if allow_directml:
|
||||||
|
try:
|
||||||
|
import torch_directml
|
||||||
|
if torch_directml.device_count() > 0:
|
||||||
|
return torch_directml.device(0)
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if allow_mps:
|
||||||
|
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
||||||
|
return "mps"
|
||||||
|
|
||||||
|
return "cpu"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Voice prompt combination
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def combine_voice_prompts(
|
||||||
|
audio_paths: List[str],
|
||||||
|
reference_texts: List[str],
|
||||||
|
*,
|
||||||
|
sample_rate: Optional[int] = None,
|
||||||
|
) -> Tuple[np.ndarray, str]:
|
||||||
|
"""
|
||||||
|
Combine multiple reference audio samples into one.
|
||||||
|
|
||||||
|
Loads each audio file, normalizes, concatenates, and joins texts.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
audio_paths: Paths to reference audio files.
|
||||||
|
reference_texts: Corresponding transcripts.
|
||||||
|
sample_rate: If set, resample audio to this rate during loading.
|
||||||
|
"""
|
||||||
|
combined_audio = []
|
||||||
|
|
||||||
|
for path in audio_paths:
|
||||||
|
kwargs = {"sample_rate": sample_rate} if sample_rate else {}
|
||||||
|
audio, _sr = load_audio(path, **kwargs)
|
||||||
|
audio = normalize_audio(audio)
|
||||||
|
combined_audio.append(audio)
|
||||||
|
|
||||||
|
mixed = np.concatenate(combined_audio)
|
||||||
|
mixed = normalize_audio(mixed)
|
||||||
|
combined_text = " ".join(reference_texts)
|
||||||
|
|
||||||
|
return mixed, combined_text
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Model loading progress tracking
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def model_load_progress(
|
||||||
|
model_name: str,
|
||||||
|
is_cached: bool,
|
||||||
|
filter_non_downloads: Optional[bool] = None,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Context manager for model loading with HF download progress tracking.
|
||||||
|
|
||||||
|
Handles the tqdm patching, progress_manager/task_manager lifecycle,
|
||||||
|
and error reporting that every backend duplicates.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_name: Progress tracking key (e.g. "qwen-tts-1.7B", "whisper-base").
|
||||||
|
is_cached: Whether the model is already downloaded.
|
||||||
|
filter_non_downloads: Whether to filter non-download tqdm bars.
|
||||||
|
Defaults to `is_cached`.
|
||||||
|
|
||||||
|
Yields:
|
||||||
|
The tracker context (already entered). The caller loads the model
|
||||||
|
inside the `with` block. The tqdm patch is torn down on exit.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
with model_load_progress("qwen-tts-1.7B", is_cached) as ctx:
|
||||||
|
self.model = SomeModel.from_pretrained(...)
|
||||||
|
"""
|
||||||
|
if filter_non_downloads is None:
|
||||||
|
filter_non_downloads = is_cached
|
||||||
|
|
||||||
|
progress_manager = get_progress_manager()
|
||||||
|
task_manager = get_task_manager()
|
||||||
|
|
||||||
|
progress_callback = create_hf_progress_callback(model_name, progress_manager)
|
||||||
|
tracker = HFProgressTracker(progress_callback, filter_non_downloads=filter_non_downloads)
|
||||||
|
|
||||||
|
tracker_context = tracker.patch_download()
|
||||||
|
tracker_context.__enter__()
|
||||||
|
|
||||||
|
if not is_cached:
|
||||||
|
task_manager.start_download(model_name)
|
||||||
|
progress_manager.update_progress(
|
||||||
|
model_name=model_name,
|
||||||
|
current=0,
|
||||||
|
total=0,
|
||||||
|
filename="Connecting to HuggingFace...",
|
||||||
|
status="downloading",
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
yield tracker_context
|
||||||
|
except Exception as e:
|
||||||
|
# Report error to both managers
|
||||||
|
progress_manager.mark_error(model_name, str(e))
|
||||||
|
task_manager.error_download(model_name, str(e))
|
||||||
|
raise
|
||||||
|
else:
|
||||||
|
# Only mark complete if we were tracking a download
|
||||||
|
if not is_cached:
|
||||||
|
progress_manager.mark_complete(model_name)
|
||||||
|
task_manager.complete_download(model_name)
|
||||||
|
finally:
|
||||||
|
tracker_context.__exit__(None, None, None)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Chatterbox f32 dtype patches
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def patch_chatterbox_f32(model) -> None:
|
||||||
|
"""
|
||||||
|
Patch float64 -> float32 dtype mismatches in upstream chatterbox.
|
||||||
|
|
||||||
|
librosa.load returns float64 numpy arrays. Multiple upstream code paths
|
||||||
|
convert these to torch tensors via torch.from_numpy() without casting,
|
||||||
|
then matmul against float32 model weights. This patches the two known
|
||||||
|
entry points:
|
||||||
|
|
||||||
|
1. S3Tokenizer.log_mel_spectrogram — audio tensor hits _mel_filters (f32)
|
||||||
|
2. VoiceEncoder.forward — float64 mel spectrograms hit LSTM weights (f32)
|
||||||
|
"""
|
||||||
|
import types
|
||||||
|
|
||||||
|
# Patch S3Tokenizer
|
||||||
|
_tokzr = model.s3gen.tokenizer
|
||||||
|
_orig_log_mel = _tokzr.log_mel_spectrogram.__func__
|
||||||
|
|
||||||
|
def _f32_log_mel(self_tokzr, audio, padding=0):
|
||||||
|
import torch as _torch
|
||||||
|
if _torch.is_tensor(audio):
|
||||||
|
audio = audio.float()
|
||||||
|
return _orig_log_mel(self_tokzr, audio, padding)
|
||||||
|
|
||||||
|
_tokzr.log_mel_spectrogram = types.MethodType(_f32_log_mel, _tokzr)
|
||||||
|
|
||||||
|
# Patch VoiceEncoder
|
||||||
|
_ve = model.ve
|
||||||
|
_orig_ve_forward = _ve.forward.__func__
|
||||||
|
|
||||||
|
def _f32_ve_forward(self_ve, mels):
|
||||||
|
return _orig_ve_forward(self_ve, mels.float())
|
||||||
|
|
||||||
|
_ve.forward = types.MethodType(_f32_ve_forward, _ve)
|
||||||
@@ -8,7 +8,6 @@ on macOS due to known MPS tensor issues.
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
import platform
|
|
||||||
import threading
|
import threading
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import ClassVar, List, Optional, Tuple
|
from typing import ClassVar, List, Optional, Tuple
|
||||||
@@ -16,9 +15,13 @@ from typing import ClassVar, List, Optional, Tuple
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
from . import TTSBackend
|
from . import TTSBackend
|
||||||
from ..utils.audio import normalize_audio, load_audio
|
from .base import (
|
||||||
from ..utils.progress import get_progress_manager
|
is_model_cached,
|
||||||
from ..utils.tasks import get_task_manager
|
get_torch_device,
|
||||||
|
combine_voice_prompts as _combine_voice_prompts,
|
||||||
|
model_load_progress,
|
||||||
|
patch_chatterbox_f32,
|
||||||
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -45,17 +48,7 @@ class ChatterboxTTSBackend:
|
|||||||
self._model_load_lock = asyncio.Lock()
|
self._model_load_lock = asyncio.Lock()
|
||||||
|
|
||||||
def _get_device(self) -> str:
|
def _get_device(self) -> str:
|
||||||
"""Get the best available device. Forces CPU on macOS (MPS issue)."""
|
return get_torch_device(force_cpu_on_mac=True)
|
||||||
if platform.system() == "Darwin":
|
|
||||||
return "cpu"
|
|
||||||
try:
|
|
||||||
import torch
|
|
||||||
|
|
||||||
if torch.cuda.is_available():
|
|
||||||
return "cuda"
|
|
||||||
except ImportError:
|
|
||||||
pass
|
|
||||||
return "cpu"
|
|
||||||
|
|
||||||
def is_loaded(self) -> bool:
|
def is_loaded(self) -> bool:
|
||||||
return self.model is not None
|
return self.model is not None
|
||||||
@@ -64,33 +57,7 @@ class ChatterboxTTSBackend:
|
|||||||
return CHATTERBOX_HF_REPO
|
return CHATTERBOX_HF_REPO
|
||||||
|
|
||||||
def _is_model_cached(self, model_size: str = "default") -> bool:
|
def _is_model_cached(self, model_size: str = "default") -> bool:
|
||||||
"""Check if the Chatterbox multilingual model is cached locally."""
|
return is_model_cached(CHATTERBOX_HF_REPO, required_files=_MTL_WEIGHT_FILES)
|
||||||
try:
|
|
||||||
from huggingface_hub import constants as hf_constants
|
|
||||||
|
|
||||||
repo_cache = Path(hf_constants.HF_HUB_CACHE) / (
|
|
||||||
"models--" + CHATTERBOX_HF_REPO.replace("/", "--")
|
|
||||||
)
|
|
||||||
|
|
||||||
if not repo_cache.exists():
|
|
||||||
return False
|
|
||||||
|
|
||||||
blobs_dir = repo_cache / "blobs"
|
|
||||||
if blobs_dir.exists() and any(blobs_dir.glob("*.incomplete")):
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Check for multilingual weight files
|
|
||||||
snapshots_dir = repo_cache / "snapshots"
|
|
||||||
if snapshots_dir.exists():
|
|
||||||
for fname in _MTL_WEIGHT_FILES:
|
|
||||||
if not any(snapshots_dir.rglob(fname)):
|
|
||||||
return False
|
|
||||||
return True
|
|
||||||
|
|
||||||
return False
|
|
||||||
except Exception as e:
|
|
||||||
logger.warning(f"Error checking Chatterbox cache: {e}")
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def load_model(self, model_size: str = "default") -> None:
|
async def load_model(self, model_size: str = "default") -> None:
|
||||||
"""Load the Chatterbox multilingual model."""
|
"""Load the Chatterbox multilingual model."""
|
||||||
@@ -103,133 +70,45 @@ class ChatterboxTTSBackend:
|
|||||||
|
|
||||||
def _load_model_sync(self):
|
def _load_model_sync(self):
|
||||||
"""Synchronous model loading."""
|
"""Synchronous model loading."""
|
||||||
from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback
|
|
||||||
|
|
||||||
progress_manager = get_progress_manager()
|
|
||||||
task_manager = get_task_manager()
|
|
||||||
model_name = "chatterbox-tts"
|
model_name = "chatterbox-tts"
|
||||||
|
|
||||||
is_cached = self._is_model_cached()
|
is_cached = self._is_model_cached()
|
||||||
|
|
||||||
# Set up HF progress tracking (intercepts tqdm for file-level progress)
|
with model_load_progress(model_name, is_cached):
|
||||||
progress_callback = create_hf_progress_callback(model_name, progress_manager)
|
|
||||||
tracker = HFProgressTracker(progress_callback, filter_non_downloads=is_cached)
|
|
||||||
tracker_context = tracker.patch_download()
|
|
||||||
tracker_context.__enter__()
|
|
||||||
|
|
||||||
if not is_cached:
|
|
||||||
task_manager.start_download(model_name)
|
|
||||||
progress_manager.update_progress(
|
|
||||||
model_name=model_name,
|
|
||||||
current=0,
|
|
||||||
total=0,
|
|
||||||
filename="Connecting to HuggingFace...",
|
|
||||||
status="downloading",
|
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
|
||||||
device = self._get_device()
|
device = self._get_device()
|
||||||
self._device = device
|
self._device = device
|
||||||
|
|
||||||
logger.info(f"Loading Chatterbox Multilingual TTS on {device}...")
|
logger.info(f"Loading Chatterbox Multilingual TTS on {device}...")
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from chatterbox.mtl_tts import ChatterboxMultilingualTTS
|
from chatterbox.mtl_tts import ChatterboxMultilingualTTS
|
||||||
|
|
||||||
# Load into a local variable first, apply all patches, then
|
if device == "cpu":
|
||||||
# assign to self.model. This avoids leaving a half-initialised
|
_orig_torch_load = torch.load
|
||||||
# model on self.model if any patch step raises an exception.
|
|
||||||
#
|
|
||||||
# Monkey-patch torch.load for CPU loading. The model's .pt files
|
|
||||||
# were saved on CUDA; from_pretrained() doesn't pass map_location
|
|
||||||
# so loading on CPU fails without this.
|
|
||||||
try:
|
|
||||||
if device == "cpu":
|
|
||||||
_orig_torch_load = torch.load
|
|
||||||
|
|
||||||
def _patched_load(*args, **kwargs):
|
def _patched_load(*args, **kwargs):
|
||||||
kwargs.setdefault("map_location", "cpu")
|
kwargs.setdefault("map_location", "cpu")
|
||||||
return _orig_torch_load(*args, **kwargs)
|
return _orig_torch_load(*args, **kwargs)
|
||||||
|
|
||||||
with ChatterboxTTSBackend._load_lock:
|
with ChatterboxTTSBackend._load_lock:
|
||||||
torch.load = _patched_load
|
torch.load = _patched_load
|
||||||
try:
|
try:
|
||||||
model = ChatterboxMultilingualTTS.from_pretrained(
|
model = ChatterboxMultilingualTTS.from_pretrained(device=device)
|
||||||
device=device,
|
finally:
|
||||||
)
|
torch.load = _orig_torch_load
|
||||||
finally:
|
else:
|
||||||
torch.load = _orig_torch_load
|
model = ChatterboxMultilingualTTS.from_pretrained(device=device)
|
||||||
else:
|
|
||||||
model = ChatterboxMultilingualTTS.from_pretrained(
|
|
||||||
device=device,
|
|
||||||
)
|
|
||||||
finally:
|
|
||||||
tracker_context.__exit__(None, None, None)
|
|
||||||
|
|
||||||
# Fix: transformers >= 4.36 defaults LlamaModel to sdpa attention
|
# Fix sdpa attention for output_attentions support
|
||||||
# which doesn't support output_attentions=True (needed by
|
|
||||||
# Chatterbox's AlignmentStreamAnalyzer). Force eager attention.
|
|
||||||
t3_tfmr = model.t3.tfmr
|
t3_tfmr = model.t3.tfmr
|
||||||
if hasattr(t3_tfmr, "config") and hasattr(
|
if hasattr(t3_tfmr, "config") and hasattr(t3_tfmr.config, "_attn_implementation"):
|
||||||
t3_tfmr.config, "_attn_implementation"
|
|
||||||
):
|
|
||||||
t3_tfmr.config._attn_implementation = "eager"
|
t3_tfmr.config._attn_implementation = "eager"
|
||||||
for layer in getattr(t3_tfmr, "layers", []):
|
for layer in getattr(t3_tfmr, "layers", []):
|
||||||
if hasattr(layer, "self_attn"):
|
if hasattr(layer, "self_attn"):
|
||||||
layer.self_attn._attn_implementation = "eager"
|
layer.self_attn._attn_implementation = "eager"
|
||||||
|
|
||||||
if not is_cached:
|
patch_chatterbox_f32(model)
|
||||||
progress_manager.mark_complete(model_name)
|
|
||||||
task_manager.complete_download(model_name)
|
|
||||||
|
|
||||||
# Patch float64 → float32 dtype mismatches in upstream chatterbox.
|
|
||||||
# librosa.load returns float64 numpy; multiple upstream code paths
|
|
||||||
# convert it to a torch tensor via torch.from_numpy() without
|
|
||||||
# casting, then matmul it against float32 model weights.
|
|
||||||
import types
|
|
||||||
|
|
||||||
# Patch S3Tokenizer (used by s3gen.tokenizer)
|
|
||||||
_tokzr = model.s3gen.tokenizer
|
|
||||||
_orig_log_mel = _tokzr.log_mel_spectrogram.__func__
|
|
||||||
|
|
||||||
def _f32_log_mel(self_tokzr, audio, padding=0):
|
|
||||||
import torch as _torch
|
|
||||||
if _torch.is_tensor(audio):
|
|
||||||
audio = audio.float()
|
|
||||||
return _orig_log_mel(self_tokzr, audio, padding)
|
|
||||||
|
|
||||||
_tokzr.log_mel_spectrogram = types.MethodType(_f32_log_mel, _tokzr)
|
|
||||||
|
|
||||||
# Patch VoiceEncoder
|
|
||||||
_ve = model.ve
|
|
||||||
_orig_ve_forward = _ve.forward.__func__
|
|
||||||
|
|
||||||
def _f32_ve_forward(self_ve, mels):
|
|
||||||
return _orig_ve_forward(self_ve, mels.float())
|
|
||||||
|
|
||||||
_ve.forward = types.MethodType(_f32_ve_forward, _ve)
|
|
||||||
|
|
||||||
# All patches applied successfully — publish the model
|
|
||||||
self.model = model
|
self.model = model
|
||||||
|
|
||||||
logger.info("Chatterbox Multilingual TTS loaded successfully")
|
logger.info("Chatterbox Multilingual TTS loaded successfully")
|
||||||
|
|
||||||
except ImportError as e:
|
|
||||||
logger.error(
|
|
||||||
"chatterbox-tts package not found. "
|
|
||||||
"Install with: pip install chatterbox-tts"
|
|
||||||
)
|
|
||||||
if not is_cached:
|
|
||||||
progress_manager.mark_error(model_name, str(e))
|
|
||||||
task_manager.error_download(model_name, str(e))
|
|
||||||
raise
|
|
||||||
except Exception as e:
|
|
||||||
import traceback
|
|
||||||
logger.error(f"Failed to load Chatterbox: {e}\n{traceback.format_exc()}")
|
|
||||||
if not is_cached:
|
|
||||||
progress_manager.mark_error(model_name, str(e))
|
|
||||||
task_manager.error_download(model_name, str(e))
|
|
||||||
raise
|
|
||||||
|
|
||||||
def unload_model(self) -> None:
|
def unload_model(self) -> None:
|
||||||
"""Unload model to free memory."""
|
"""Unload model to free memory."""
|
||||||
@@ -268,17 +147,7 @@ class ChatterboxTTSBackend:
|
|||||||
audio_paths: List[str],
|
audio_paths: List[str],
|
||||||
reference_texts: List[str],
|
reference_texts: List[str],
|
||||||
) -> Tuple[np.ndarray, str]:
|
) -> Tuple[np.ndarray, str]:
|
||||||
"""Combine multiple reference samples."""
|
return await _combine_voice_prompts(audio_paths, reference_texts)
|
||||||
combined_audio = []
|
|
||||||
for path in audio_paths:
|
|
||||||
audio, _sr = load_audio(path)
|
|
||||||
audio = normalize_audio(audio)
|
|
||||||
combined_audio.append(audio)
|
|
||||||
|
|
||||||
mixed = np.concatenate(combined_audio)
|
|
||||||
mixed = normalize_audio(mixed)
|
|
||||||
combined_text = " ".join(reference_texts)
|
|
||||||
return mixed, combined_text
|
|
||||||
|
|
||||||
# Per-language generation defaults. Lower temp + higher cfg = clearer speech.
|
# Per-language generation defaults. Lower temp + higher cfg = clearer speech.
|
||||||
_LANG_DEFAULTS: ClassVar[dict] = {
|
_LANG_DEFAULTS: ClassVar[dict] = {
|
||||||
|
|||||||
@@ -8,7 +8,6 @@ Forces CPU on macOS due to known MPS tensor issues.
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
import platform
|
|
||||||
import threading
|
import threading
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import ClassVar, List, Optional, Tuple
|
from typing import ClassVar, List, Optional, Tuple
|
||||||
@@ -16,9 +15,13 @@ from typing import ClassVar, List, Optional, Tuple
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
from . import TTSBackend
|
from . import TTSBackend
|
||||||
from ..utils.audio import normalize_audio, load_audio
|
from .base import (
|
||||||
from ..utils.progress import get_progress_manager
|
is_model_cached,
|
||||||
from ..utils.tasks import get_task_manager
|
get_torch_device,
|
||||||
|
combine_voice_prompts as _combine_voice_prompts,
|
||||||
|
model_load_progress,
|
||||||
|
patch_chatterbox_f32,
|
||||||
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -45,17 +48,7 @@ class ChatterboxTurboTTSBackend:
|
|||||||
self._model_load_lock = asyncio.Lock()
|
self._model_load_lock = asyncio.Lock()
|
||||||
|
|
||||||
def _get_device(self) -> str:
|
def _get_device(self) -> str:
|
||||||
"""Get the best available device. Forces CPU on macOS (MPS issue)."""
|
return get_torch_device(force_cpu_on_mac=True)
|
||||||
if platform.system() == "Darwin":
|
|
||||||
return "cpu"
|
|
||||||
try:
|
|
||||||
import torch
|
|
||||||
|
|
||||||
if torch.cuda.is_available():
|
|
||||||
return "cuda"
|
|
||||||
except ImportError:
|
|
||||||
pass
|
|
||||||
return "cpu"
|
|
||||||
|
|
||||||
def is_loaded(self) -> bool:
|
def is_loaded(self) -> bool:
|
||||||
return self.model is not None
|
return self.model is not None
|
||||||
@@ -64,33 +57,7 @@ class ChatterboxTurboTTSBackend:
|
|||||||
return CHATTERBOX_TURBO_HF_REPO
|
return CHATTERBOX_TURBO_HF_REPO
|
||||||
|
|
||||||
def _is_model_cached(self, model_size: str = "default") -> bool:
|
def _is_model_cached(self, model_size: str = "default") -> bool:
|
||||||
"""Check if the Chatterbox Turbo model is cached locally."""
|
return is_model_cached(CHATTERBOX_TURBO_HF_REPO, required_files=_TURBO_WEIGHT_FILES)
|
||||||
try:
|
|
||||||
from huggingface_hub import constants as hf_constants
|
|
||||||
|
|
||||||
repo_cache = Path(hf_constants.HF_HUB_CACHE) / (
|
|
||||||
"models--" + CHATTERBOX_TURBO_HF_REPO.replace("/", "--")
|
|
||||||
)
|
|
||||||
|
|
||||||
if not repo_cache.exists():
|
|
||||||
return False
|
|
||||||
|
|
||||||
blobs_dir = repo_cache / "blobs"
|
|
||||||
if blobs_dir.exists() and any(blobs_dir.glob("*.incomplete")):
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Check for turbo weight files
|
|
||||||
snapshots_dir = repo_cache / "snapshots"
|
|
||||||
if snapshots_dir.exists():
|
|
||||||
for fname in _TURBO_WEIGHT_FILES:
|
|
||||||
if not any(snapshots_dir.rglob(fname)):
|
|
||||||
return False
|
|
||||||
return True
|
|
||||||
|
|
||||||
return False
|
|
||||||
except Exception as e:
|
|
||||||
logger.warning(f"Error checking Chatterbox Turbo cache: {e}")
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def load_model(self, model_size: str = "default") -> None:
|
async def load_model(self, model_size: str = "default") -> None:
|
||||||
"""Load the Chatterbox Turbo model."""
|
"""Load the Chatterbox Turbo model."""
|
||||||
@@ -103,59 +70,24 @@ class ChatterboxTurboTTSBackend:
|
|||||||
|
|
||||||
def _load_model_sync(self):
|
def _load_model_sync(self):
|
||||||
"""Synchronous model loading."""
|
"""Synchronous model loading."""
|
||||||
from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback
|
|
||||||
|
|
||||||
progress_manager = get_progress_manager()
|
|
||||||
task_manager = get_task_manager()
|
|
||||||
model_name = "chatterbox-turbo"
|
model_name = "chatterbox-turbo"
|
||||||
|
|
||||||
is_cached = self._is_model_cached()
|
is_cached = self._is_model_cached()
|
||||||
|
|
||||||
# Set up HF progress tracking (intercepts tqdm for file-level progress)
|
with model_load_progress(model_name, is_cached):
|
||||||
progress_callback = create_hf_progress_callback(model_name, progress_manager)
|
|
||||||
tracker = HFProgressTracker(progress_callback, filter_non_downloads=is_cached)
|
|
||||||
tracker_context = tracker.patch_download()
|
|
||||||
tracker_context.__enter__()
|
|
||||||
|
|
||||||
if not is_cached:
|
|
||||||
task_manager.start_download(model_name)
|
|
||||||
progress_manager.update_progress(
|
|
||||||
model_name=model_name,
|
|
||||||
current=0,
|
|
||||||
total=0,
|
|
||||||
filename="Connecting to HuggingFace...",
|
|
||||||
status="downloading",
|
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
|
||||||
device = self._get_device()
|
device = self._get_device()
|
||||||
self._device = device
|
self._device = device
|
||||||
|
|
||||||
logger.info(f"Loading Chatterbox Turbo TTS on {device}...")
|
logger.info(f"Loading Chatterbox Turbo TTS on {device}...")
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from huggingface_hub import snapshot_download
|
from huggingface_hub import snapshot_download
|
||||||
from chatterbox.tts_turbo import ChatterboxTurboTTS
|
from chatterbox.tts_turbo import ChatterboxTurboTTS
|
||||||
|
|
||||||
# Download model files ourselves so we can pass token=None
|
local_path = snapshot_download(
|
||||||
# (upstream from_pretrained passes token=True which requires
|
repo_id=CHATTERBOX_TURBO_HF_REPO,
|
||||||
# a stored HF token even though the repo is public).
|
token=None,
|
||||||
try:
|
allow_patterns=["*.safetensors", "*.json", "*.txt", "*.pt", "*.model"],
|
||||||
local_path = snapshot_download(
|
)
|
||||||
repo_id=CHATTERBOX_TURBO_HF_REPO,
|
|
||||||
token=None,
|
|
||||||
allow_patterns=[
|
|
||||||
"*.safetensors", "*.json", "*.txt", "*.pt", "*.model",
|
|
||||||
],
|
|
||||||
)
|
|
||||||
finally:
|
|
||||||
tracker_context.__exit__(None, None, None)
|
|
||||||
|
|
||||||
# Monkey-patch torch.load for CPU loading. The model's .pt files
|
|
||||||
# were saved on CUDA; from_local() doesn't pass map_location
|
|
||||||
# so loading on CPU fails without this.
|
|
||||||
# Load into a local var, apply patches, then publish to
|
|
||||||
# self.model so a failed patch doesn't leave us half-initialised.
|
|
||||||
if device == "cpu":
|
if device == "cpu":
|
||||||
_orig_torch_load = torch.load
|
_orig_torch_load = torch.load
|
||||||
|
|
||||||
@@ -166,74 +98,16 @@ class ChatterboxTurboTTSBackend:
|
|||||||
with ChatterboxTurboTTSBackend._load_lock:
|
with ChatterboxTurboTTSBackend._load_lock:
|
||||||
torch.load = _patched_load
|
torch.load = _patched_load
|
||||||
try:
|
try:
|
||||||
model = ChatterboxTurboTTS.from_local(
|
model = ChatterboxTurboTTS.from_local(local_path, device)
|
||||||
local_path, device,
|
|
||||||
)
|
|
||||||
finally:
|
finally:
|
||||||
torch.load = _orig_torch_load
|
torch.load = _orig_torch_load
|
||||||
else:
|
else:
|
||||||
model = ChatterboxTurboTTS.from_local(
|
model = ChatterboxTurboTTS.from_local(local_path, device)
|
||||||
local_path, device,
|
|
||||||
)
|
|
||||||
|
|
||||||
if not is_cached:
|
patch_chatterbox_f32(model)
|
||||||
progress_manager.mark_complete(model_name)
|
|
||||||
task_manager.complete_download(model_name)
|
|
||||||
|
|
||||||
# Patch float64 → float32 dtype mismatches in upstream chatterbox.
|
|
||||||
# librosa.load returns float64 numpy; multiple upstream code paths
|
|
||||||
# convert it to a torch tensor via torch.from_numpy() without
|
|
||||||
# casting, then matmul it against float32 model weights.
|
|
||||||
# We patch the two known entry points:
|
|
||||||
#
|
|
||||||
# 1. S3Tokenizer.log_mel_spectrogram — the audio tensor from
|
|
||||||
# librosa hits _mel_filters (float32) in a matmul.
|
|
||||||
# 2. VoiceEncoder.forward — float64 mel spectrograms hit the
|
|
||||||
# float32 LSTM weights.
|
|
||||||
import types
|
|
||||||
|
|
||||||
# Patch S3Tokenizer (used by s3gen.tokenizer)
|
|
||||||
_tokzr = model.s3gen.tokenizer
|
|
||||||
_orig_log_mel = _tokzr.log_mel_spectrogram.__func__
|
|
||||||
|
|
||||||
def _f32_log_mel(self_tokzr, audio, padding=0):
|
|
||||||
import torch as _torch
|
|
||||||
if _torch.is_tensor(audio):
|
|
||||||
audio = audio.float()
|
|
||||||
return _orig_log_mel(self_tokzr, audio, padding)
|
|
||||||
|
|
||||||
_tokzr.log_mel_spectrogram = types.MethodType(_f32_log_mel, _tokzr)
|
|
||||||
|
|
||||||
# Patch VoiceEncoder
|
|
||||||
_ve = model.ve
|
|
||||||
_orig_ve_forward = _ve.forward.__func__
|
|
||||||
|
|
||||||
def _f32_ve_forward(self_ve, mels):
|
|
||||||
return _orig_ve_forward(self_ve, mels.float())
|
|
||||||
|
|
||||||
_ve.forward = types.MethodType(_f32_ve_forward, _ve)
|
|
||||||
|
|
||||||
# Only publish after all patches succeed
|
|
||||||
self.model = model
|
self.model = model
|
||||||
|
|
||||||
logger.info("Chatterbox Turbo TTS loaded successfully")
|
logger.info("Chatterbox Turbo TTS loaded successfully")
|
||||||
|
|
||||||
except ImportError as e:
|
|
||||||
logger.error(
|
|
||||||
"chatterbox-tts package not found. "
|
|
||||||
"Install with: pip install chatterbox-tts"
|
|
||||||
)
|
|
||||||
if not is_cached:
|
|
||||||
progress_manager.mark_error(model_name, str(e))
|
|
||||||
task_manager.error_download(model_name, str(e))
|
|
||||||
raise
|
|
||||||
except Exception as e:
|
|
||||||
import traceback
|
|
||||||
logger.error(f"Failed to load Chatterbox Turbo: {e}\n{traceback.format_exc()}")
|
|
||||||
if not is_cached:
|
|
||||||
progress_manager.mark_error(model_name, str(e))
|
|
||||||
task_manager.error_download(model_name, str(e))
|
|
||||||
raise
|
|
||||||
|
|
||||||
def unload_model(self) -> None:
|
def unload_model(self) -> None:
|
||||||
"""Unload model to free memory."""
|
"""Unload model to free memory."""
|
||||||
@@ -271,17 +145,7 @@ class ChatterboxTurboTTSBackend:
|
|||||||
audio_paths: List[str],
|
audio_paths: List[str],
|
||||||
reference_texts: List[str],
|
reference_texts: List[str],
|
||||||
) -> Tuple[np.ndarray, str]:
|
) -> Tuple[np.ndarray, str]:
|
||||||
"""Combine multiple reference samples."""
|
return await _combine_voice_prompts(audio_paths, reference_texts)
|
||||||
combined_audio = []
|
|
||||||
for path in audio_paths:
|
|
||||||
audio, _sr = load_audio(path)
|
|
||||||
audio = normalize_audio(audio)
|
|
||||||
combined_audio.append(audio)
|
|
||||||
|
|
||||||
mixed = np.concatenate(combined_audio)
|
|
||||||
mixed = normalize_audio(mixed)
|
|
||||||
combined_text = " ".join(reference_texts)
|
|
||||||
return mixed, combined_text
|
|
||||||
|
|
||||||
async def generate(
|
async def generate(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -7,16 +7,13 @@ Wraps the LuxTTS (ZipVoice) model for zero-shot voice cloning.
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
from pathlib import Path
|
from typing import Optional, Tuple
|
||||||
from typing import List, Optional, Tuple
|
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
from . import TTSBackend
|
from . import TTSBackend
|
||||||
from ..utils.audio import normalize_audio, load_audio
|
from .base import is_model_cached, get_torch_device, combine_voice_prompts as _combine_voice_prompts, model_load_progress
|
||||||
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
|
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
|
||||||
from ..utils.progress import get_progress_manager
|
|
||||||
from ..utils.tasks import get_task_manager
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -33,14 +30,7 @@ class LuxTTSBackend:
|
|||||||
self._device = None
|
self._device = None
|
||||||
|
|
||||||
def _get_device(self) -> str:
|
def _get_device(self) -> str:
|
||||||
"""Get the best available device."""
|
return get_torch_device(allow_mps=True)
|
||||||
import torch
|
|
||||||
|
|
||||||
if torch.cuda.is_available():
|
|
||||||
return "cuda"
|
|
||||||
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
|
||||||
return "mps"
|
|
||||||
return "cpu"
|
|
||||||
|
|
||||||
def is_loaded(self) -> bool:
|
def is_loaded(self) -> bool:
|
||||||
return self.model is not None
|
return self.model is not None
|
||||||
@@ -55,35 +45,10 @@ class LuxTTSBackend:
|
|||||||
return LUXTTS_HF_REPO
|
return LUXTTS_HF_REPO
|
||||||
|
|
||||||
def _is_model_cached(self, model_size: str = "default") -> bool:
|
def _is_model_cached(self, model_size: str = "default") -> bool:
|
||||||
"""Check if LuxTTS model weights are cached locally."""
|
return is_model_cached(
|
||||||
try:
|
LUXTTS_HF_REPO,
|
||||||
from huggingface_hub import constants as hf_constants
|
weight_extensions=(".pt", ".safetensors", ".onnx", ".bin"),
|
||||||
|
)
|
||||||
repo_cache = (
|
|
||||||
Path(hf_constants.HF_HUB_CACHE)
|
|
||||||
/ ("models--" + LUXTTS_HF_REPO.replace("/", "--"))
|
|
||||||
)
|
|
||||||
|
|
||||||
if not repo_cache.exists():
|
|
||||||
return False
|
|
||||||
|
|
||||||
blobs_dir = repo_cache / "blobs"
|
|
||||||
if blobs_dir.exists() and any(blobs_dir.glob("*.incomplete")):
|
|
||||||
return False
|
|
||||||
|
|
||||||
snapshots_dir = repo_cache / "snapshots"
|
|
||||||
if snapshots_dir.exists():
|
|
||||||
has_weights = any(snapshots_dir.rglob("*.pt")) or any(
|
|
||||||
snapshots_dir.rglob("*.safetensors")
|
|
||||||
) or any(snapshots_dir.rglob("*.onnx")) or any(
|
|
||||||
snapshots_dir.rglob("*.bin")
|
|
||||||
)
|
|
||||||
return has_weights
|
|
||||||
|
|
||||||
return False
|
|
||||||
except Exception as e:
|
|
||||||
logger.warning(f"Error checking LuxTTS cache: {e}")
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def load_model(self, model_size: str = "default") -> None:
|
async def load_model(self, model_size: str = "default") -> None:
|
||||||
"""Load the LuxTTS model."""
|
"""Load the LuxTTS model."""
|
||||||
@@ -93,68 +58,25 @@ class LuxTTSBackend:
|
|||||||
await asyncio.to_thread(self._load_model_sync)
|
await asyncio.to_thread(self._load_model_sync)
|
||||||
|
|
||||||
def _load_model_sync(self):
|
def _load_model_sync(self):
|
||||||
"""Synchronous model loading."""
|
|
||||||
from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback
|
|
||||||
|
|
||||||
progress_manager = get_progress_manager()
|
|
||||||
task_manager = get_task_manager()
|
|
||||||
model_name = "luxtts"
|
model_name = "luxtts"
|
||||||
|
|
||||||
is_cached = self._is_model_cached()
|
is_cached = self._is_model_cached()
|
||||||
|
|
||||||
# Set up HF progress tracking (intercepts tqdm for file-level progress)
|
with model_load_progress(model_name, is_cached):
|
||||||
progress_callback = create_hf_progress_callback(model_name, progress_manager)
|
|
||||||
tracker = HFProgressTracker(progress_callback, filter_non_downloads=is_cached)
|
|
||||||
tracker_context = tracker.patch_download()
|
|
||||||
tracker_context.__enter__()
|
|
||||||
|
|
||||||
if not is_cached:
|
|
||||||
task_manager.start_download(model_name)
|
|
||||||
progress_manager.update_progress(
|
|
||||||
model_name=model_name,
|
|
||||||
current=0,
|
|
||||||
total=0,
|
|
||||||
filename="Connecting to HuggingFace...",
|
|
||||||
status="downloading",
|
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
|
||||||
from zipvoice.luxvoice import LuxTTS
|
from zipvoice.luxvoice import LuxTTS
|
||||||
|
|
||||||
device = self.device
|
device = self.device
|
||||||
logger.info(f"Loading LuxTTS on {device}...")
|
logger.info(f"Loading LuxTTS on {device}...")
|
||||||
|
|
||||||
# LuxTTS constructor downloads model and loads everything
|
if device == "cpu":
|
||||||
try:
|
import os
|
||||||
if device == "cpu":
|
threads = os.cpu_count() or 4
|
||||||
import os
|
self.model = LuxTTS(
|
||||||
threads = os.cpu_count() or 4
|
model_path=LUXTTS_HF_REPO, device="cpu", threads=min(threads, 8),
|
||||||
self.model = LuxTTS(
|
)
|
||||||
model_path=LUXTTS_HF_REPO,
|
else:
|
||||||
device="cpu",
|
self.model = LuxTTS(model_path=LUXTTS_HF_REPO, device=device)
|
||||||
threads=min(threads, 8),
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
self.model = LuxTTS(
|
|
||||||
model_path=LUXTTS_HF_REPO,
|
|
||||||
device=device,
|
|
||||||
)
|
|
||||||
finally:
|
|
||||||
tracker_context.__exit__(None, None, None)
|
|
||||||
|
|
||||||
if not is_cached:
|
logger.info("LuxTTS loaded successfully")
|
||||||
progress_manager.mark_complete(model_name)
|
|
||||||
task_manager.complete_download(model_name)
|
|
||||||
|
|
||||||
logger.info("LuxTTS loaded successfully")
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
import traceback
|
|
||||||
logger.error(f"Failed to load LuxTTS: {e}\n{traceback.format_exc()}")
|
|
||||||
if not is_cached:
|
|
||||||
progress_manager.mark_error(model_name, str(e))
|
|
||||||
task_manager.error_download(model_name, str(e))
|
|
||||||
raise
|
|
||||||
|
|
||||||
def unload_model(self) -> None:
|
def unload_model(self) -> None:
|
||||||
"""Unload model to free memory."""
|
"""Unload model to free memory."""
|
||||||
@@ -205,28 +127,8 @@ class LuxTTSBackend:
|
|||||||
|
|
||||||
return encoded, False
|
return encoded, False
|
||||||
|
|
||||||
async def combine_voice_prompts(
|
async def combine_voice_prompts(self, audio_paths, reference_texts):
|
||||||
self,
|
return await _combine_voice_prompts(audio_paths, reference_texts, sample_rate=24000)
|
||||||
audio_paths: List[str],
|
|
||||||
reference_texts: List[str],
|
|
||||||
) -> Tuple[np.ndarray, str]:
|
|
||||||
"""
|
|
||||||
Combine multiple reference samples.
|
|
||||||
|
|
||||||
LuxTTS doesn't have native multi-prompt support, so we concatenate
|
|
||||||
the audio and let encode_prompt handle the combined clip.
|
|
||||||
"""
|
|
||||||
combined_audio = []
|
|
||||||
for path in audio_paths:
|
|
||||||
audio, _sr = load_audio(path, sample_rate=24000)
|
|
||||||
audio = normalize_audio(audio)
|
|
||||||
combined_audio.append(audio)
|
|
||||||
|
|
||||||
mixed = np.concatenate(combined_audio)
|
|
||||||
mixed = normalize_audio(mixed)
|
|
||||||
combined_text = " ".join(reference_texts)
|
|
||||||
|
|
||||||
return mixed, combined_text
|
|
||||||
|
|
||||||
async def generate(
|
async def generate(
|
||||||
self,
|
self,
|
||||||
|
|||||||
+48
-289
@@ -14,18 +14,9 @@ from ..utils.hf_offline_patch import patch_huggingface_hub_offline, ensure_origi
|
|||||||
patch_huggingface_hub_offline()
|
patch_huggingface_hub_offline()
|
||||||
ensure_original_qwen_config_cached()
|
ensure_original_qwen_config_cached()
|
||||||
|
|
||||||
from . import TTSBackend, STTBackend
|
from . import TTSBackend, STTBackend, LANGUAGE_CODE_TO_NAME, WHISPER_HF_REPOS
|
||||||
|
from .base import is_model_cached, combine_voice_prompts as _combine_voice_prompts, model_load_progress
|
||||||
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
|
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
|
||||||
from ..utils.audio import normalize_audio, load_audio
|
|
||||||
from ..utils.progress import get_progress_manager
|
|
||||||
from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback
|
|
||||||
from ..utils.tasks import get_task_manager
|
|
||||||
|
|
||||||
LANGUAGE_CODE_TO_NAME = {
|
|
||||||
"zh": "chinese", "en": "english", "ja": "japanese", "ko": "korean",
|
|
||||||
"de": "german", "fr": "french", "ru": "russian", "pt": "portuguese",
|
|
||||||
"es": "spanish", "it": "italian",
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
class MLXTTSBackend:
|
class MLXTTSBackend:
|
||||||
@@ -66,45 +57,10 @@ class MLXTTSBackend:
|
|||||||
return hf_model_id
|
return hf_model_id
|
||||||
|
|
||||||
def _is_model_cached(self, model_size: str) -> bool:
|
def _is_model_cached(self, model_size: str) -> bool:
|
||||||
"""
|
return is_model_cached(
|
||||||
Check if the model is already cached locally AND fully downloaded.
|
self._get_model_path(model_size),
|
||||||
|
weight_extensions=(".safetensors", ".bin", ".npz"),
|
||||||
Args:
|
)
|
||||||
model_size: Model size to check
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True if model is fully cached, False if missing or incomplete
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
from huggingface_hub import constants as hf_constants
|
|
||||||
model_path = self._get_model_path(model_size)
|
|
||||||
repo_cache = Path(hf_constants.HF_HUB_CACHE) / ("models--" + model_path.replace("/", "--"))
|
|
||||||
|
|
||||||
if not repo_cache.exists():
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Check for .incomplete files - if any exist, download is still in progress
|
|
||||||
blobs_dir = repo_cache / "blobs"
|
|
||||||
if blobs_dir.exists() and any(blobs_dir.glob("*.incomplete")):
|
|
||||||
print(f"[_is_model_cached] Found .incomplete files for {model_size}, treating as not cached")
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Check that actual model weight files exist in snapshots
|
|
||||||
snapshots_dir = repo_cache / "snapshots"
|
|
||||||
if snapshots_dir.exists():
|
|
||||||
has_weights = (
|
|
||||||
any(snapshots_dir.rglob("*.safetensors")) or
|
|
||||||
any(snapshots_dir.rglob("*.bin")) or
|
|
||||||
any(snapshots_dir.rglob("*.npz"))
|
|
||||||
)
|
|
||||||
if not has_weights:
|
|
||||||
print(f"[_is_model_cached] No model weights found for {model_size}, treating as not cached")
|
|
||||||
return False
|
|
||||||
|
|
||||||
return True
|
|
||||||
except Exception as e:
|
|
||||||
print(f"[_is_model_cached] Error checking cache for {model_size}: {e}")
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def load_model_async(self, model_size: Optional[str] = None):
|
async def load_model_async(self, model_size: Optional[str] = None):
|
||||||
"""
|
"""
|
||||||
@@ -132,102 +88,39 @@ class MLXTTSBackend:
|
|||||||
|
|
||||||
def _load_model_sync(self, model_size: str):
|
def _load_model_sync(self, model_size: str):
|
||||||
"""Synchronous model loading."""
|
"""Synchronous model loading."""
|
||||||
|
model_path = self._get_model_path(model_size)
|
||||||
|
model_name = f"qwen-tts-{model_size}"
|
||||||
|
is_cached = self._is_model_cached(model_size)
|
||||||
|
|
||||||
|
# Force offline mode when cached to avoid network requests
|
||||||
|
original_hf_hub_offline = os.environ.get("HF_HUB_OFFLINE")
|
||||||
|
if is_cached:
|
||||||
|
os.environ["HF_HUB_OFFLINE"] = "1"
|
||||||
|
print(f"[PATCH] Model {model_size} is cached, forcing HF_HUB_OFFLINE=1 to avoid network requests")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Get model path BEFORE importing mlx_audio
|
with model_load_progress(model_name, is_cached):
|
||||||
model_path = self._get_model_path(model_size)
|
from mlx_audio.tts import load
|
||||||
|
print(f"Loading MLX TTS model {model_size}...")
|
||||||
# Set up progress tracking
|
|
||||||
progress_manager = get_progress_manager()
|
|
||||||
task_manager = get_task_manager()
|
|
||||||
model_name = f"qwen-tts-{model_size}"
|
|
||||||
|
|
||||||
# Check if model is already cached
|
|
||||||
is_cached = self._is_model_cached(model_size)
|
|
||||||
|
|
||||||
# Set up progress callback
|
|
||||||
# If cached: filter out non-download progress
|
|
||||||
# If not cached: report all progress (we're actually downloading)
|
|
||||||
progress_callback = create_hf_progress_callback(model_name, progress_manager)
|
|
||||||
tracker = HFProgressTracker(progress_callback, filter_non_downloads=is_cached)
|
|
||||||
|
|
||||||
print(f"Loading MLX TTS model {model_size}...")
|
|
||||||
|
|
||||||
# Only track download progress if model is NOT cached
|
|
||||||
if not is_cached:
|
|
||||||
# Start tracking download task
|
|
||||||
task_manager.start_download(model_name)
|
|
||||||
|
|
||||||
# Initialize progress state so SSE endpoint has initial data to send
|
try:
|
||||||
# This provides immediate feedback while HuggingFace fetches metadata
|
|
||||||
progress_manager.update_progress(
|
|
||||||
model_name=model_name,
|
|
||||||
current=0,
|
|
||||||
total=0, # Will be updated once actual total is known
|
|
||||||
filename="Connecting to HuggingFace...",
|
|
||||||
status="downloading",
|
|
||||||
)
|
|
||||||
|
|
||||||
# IMPORTANT: Patch tqdm BEFORE importing mlx_audio
|
|
||||||
# Otherwise mlx_audio caches reference to original tqdm
|
|
||||||
tracker_context = tracker.patch_download()
|
|
||||||
tracker_context.__enter__()
|
|
||||||
|
|
||||||
# PATCH: Force offline mode when model is already cached
|
|
||||||
# This prevents crashes when HuggingFace is unreachable
|
|
||||||
original_hf_hub_offline = os.environ.get("HF_HUB_OFFLINE")
|
|
||||||
if is_cached:
|
|
||||||
os.environ["HF_HUB_OFFLINE"] = "1"
|
|
||||||
print(f"[PATCH] Model {model_size} is cached, forcing HF_HUB_OFFLINE=1 to avoid network requests")
|
|
||||||
|
|
||||||
# Import mlx_audio AFTER patching tqdm
|
|
||||||
from mlx_audio.tts import load
|
|
||||||
|
|
||||||
# Load MLX model (downloads automatically)
|
|
||||||
try:
|
|
||||||
self.model = load(model_path)
|
|
||||||
except Exception as load_error:
|
|
||||||
# If offline mode failed, try with network enabled as fallback
|
|
||||||
if is_cached and "offline" in str(load_error).lower():
|
|
||||||
print(f"[PATCH] Offline load failed, trying with network: {load_error}")
|
|
||||||
os.environ.pop("HF_HUB_OFFLINE", None)
|
|
||||||
self.model = load(model_path)
|
self.model = load(model_path)
|
||||||
else:
|
except Exception as load_error:
|
||||||
raise
|
if is_cached and "offline" in str(load_error).lower():
|
||||||
finally:
|
print(f"[PATCH] Offline load failed, trying with network: {load_error}")
|
||||||
# Exit the patch context
|
os.environ.pop("HF_HUB_OFFLINE", None)
|
||||||
tracker_context.__exit__(None, None, None)
|
self.model = load(model_path)
|
||||||
# Restore original HF_HUB_OFFLINE setting
|
else:
|
||||||
if original_hf_hub_offline is not None:
|
raise
|
||||||
os.environ["HF_HUB_OFFLINE"] = original_hf_hub_offline
|
finally:
|
||||||
else:
|
if original_hf_hub_offline is not None:
|
||||||
os.environ.pop("HF_HUB_OFFLINE", None)
|
os.environ["HF_HUB_OFFLINE"] = original_hf_hub_offline
|
||||||
|
else:
|
||||||
# Only mark download as complete if we were tracking it
|
os.environ.pop("HF_HUB_OFFLINE", None)
|
||||||
if not is_cached:
|
|
||||||
progress_manager.mark_complete(model_name)
|
self._current_model_size = model_size
|
||||||
task_manager.complete_download(model_name)
|
self.model_size = model_size
|
||||||
|
print(f"MLX TTS model {model_size} loaded successfully")
|
||||||
self._current_model_size = model_size
|
|
||||||
self.model_size = model_size
|
|
||||||
|
|
||||||
print(f"MLX TTS model {model_size} loaded successfully")
|
|
||||||
|
|
||||||
except ImportError as e:
|
|
||||||
print(f"Error: mlx_audio package not found. Install with: pip install mlx-audio")
|
|
||||||
progress_manager = get_progress_manager()
|
|
||||||
task_manager = get_task_manager()
|
|
||||||
model_name = f"qwen-tts-{model_size}"
|
|
||||||
progress_manager.mark_error(model_name, str(e))
|
|
||||||
task_manager.error_download(model_name, str(e))
|
|
||||||
raise
|
|
||||||
except Exception as e:
|
|
||||||
print(f"Error loading MLX TTS model: {e}")
|
|
||||||
progress_manager = get_progress_manager()
|
|
||||||
task_manager = get_task_manager()
|
|
||||||
model_name = f"qwen-tts-{model_size}"
|
|
||||||
progress_manager.mark_error(model_name, str(e))
|
|
||||||
task_manager.error_download(model_name, str(e))
|
|
||||||
raise
|
|
||||||
|
|
||||||
def unload_model(self):
|
def unload_model(self):
|
||||||
"""Unload the model to free memory."""
|
"""Unload the model to free memory."""
|
||||||
@@ -288,36 +181,8 @@ class MLXTTSBackend:
|
|||||||
|
|
||||||
return voice_prompt_items, False
|
return voice_prompt_items, False
|
||||||
|
|
||||||
async def combine_voice_prompts(
|
async def combine_voice_prompts(self, audio_paths, reference_texts):
|
||||||
self,
|
return await _combine_voice_prompts(audio_paths, reference_texts)
|
||||||
audio_paths: List[str],
|
|
||||||
reference_texts: List[str],
|
|
||||||
) -> Tuple[np.ndarray, str]:
|
|
||||||
"""
|
|
||||||
Combine multiple reference samples for better quality.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
audio_paths: List of audio file paths
|
|
||||||
reference_texts: List of reference texts
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tuple of (combined_audio, combined_text)
|
|
||||||
"""
|
|
||||||
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 generate(
|
async def generate(
|
||||||
self,
|
self,
|
||||||
@@ -413,14 +278,6 @@ class MLXTTSBackend:
|
|||||||
return audio, sample_rate
|
return audio, sample_rate
|
||||||
|
|
||||||
|
|
||||||
WHISPER_HF_REPOS = {
|
|
||||||
"base": "openai/whisper-base",
|
|
||||||
"small": "openai/whisper-small",
|
|
||||||
"medium": "openai/whisper-medium",
|
|
||||||
"large": "openai/whisper-large-v3",
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
class MLXSTTBackend:
|
class MLXSTTBackend:
|
||||||
"""MLX-based STT backend using mlx-audio Whisper."""
|
"""MLX-based STT backend using mlx-audio Whisper."""
|
||||||
|
|
||||||
@@ -433,45 +290,8 @@ class MLXSTTBackend:
|
|||||||
return self.model is not None
|
return self.model is not None
|
||||||
|
|
||||||
def _is_model_cached(self, model_size: str) -> bool:
|
def _is_model_cached(self, model_size: str) -> bool:
|
||||||
"""
|
hf_repo = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}")
|
||||||
Check if the Whisper model is already cached locally AND fully downloaded.
|
return is_model_cached(hf_repo, weight_extensions=(".safetensors", ".bin", ".npz"))
|
||||||
|
|
||||||
Args:
|
|
||||||
model_size: Model size to check
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True if model is fully cached, False if missing or incomplete
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
from huggingface_hub import constants as hf_constants
|
|
||||||
hf_repo = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}")
|
|
||||||
repo_cache = Path(hf_constants.HF_HUB_CACHE) / ("models--" + hf_repo.replace("/", "--"))
|
|
||||||
|
|
||||||
if not repo_cache.exists():
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Check for .incomplete files - if any exist, download is still in progress
|
|
||||||
blobs_dir = repo_cache / "blobs"
|
|
||||||
if blobs_dir.exists() and any(blobs_dir.glob("*.incomplete")):
|
|
||||||
print(f"[_is_model_cached] Found .incomplete files for whisper-{model_size}, treating as not cached")
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Check that actual model weight files exist in snapshots
|
|
||||||
snapshots_dir = repo_cache / "snapshots"
|
|
||||||
if snapshots_dir.exists():
|
|
||||||
has_weights = (
|
|
||||||
any(snapshots_dir.rglob("*.safetensors")) or
|
|
||||||
any(snapshots_dir.rglob("*.bin")) or
|
|
||||||
any(snapshots_dir.rglob("*.npz"))
|
|
||||||
)
|
|
||||||
if not has_weights:
|
|
||||||
print(f"[_is_model_cached] No model weights found for whisper-{model_size}, treating as not cached")
|
|
||||||
return False
|
|
||||||
|
|
||||||
return True
|
|
||||||
except Exception as e:
|
|
||||||
print(f"[_is_model_cached] Error checking cache for whisper-{model_size}: {e}")
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def load_model_async(self, model_size: Optional[str] = None):
|
async def load_model_async(self, model_size: Optional[str] = None):
|
||||||
"""
|
"""
|
||||||
@@ -494,78 +314,17 @@ class MLXSTTBackend:
|
|||||||
|
|
||||||
def _load_model_sync(self, model_size: str):
|
def _load_model_sync(self, model_size: str):
|
||||||
"""Synchronous model loading."""
|
"""Synchronous model loading."""
|
||||||
try:
|
progress_model_name = f"whisper-{model_size}"
|
||||||
progress_manager = get_progress_manager()
|
is_cached = self._is_model_cached(model_size)
|
||||||
task_manager = get_task_manager()
|
|
||||||
progress_model_name = f"whisper-{model_size}"
|
with model_load_progress(progress_model_name, is_cached):
|
||||||
|
|
||||||
# Check if model is already cached
|
|
||||||
is_cached = self._is_model_cached(model_size)
|
|
||||||
|
|
||||||
# Set up progress callback and tracker
|
|
||||||
# If cached: filter out non-download progress
|
|
||||||
# If not cached: report all progress (we're actually downloading)
|
|
||||||
progress_callback = create_hf_progress_callback(progress_model_name, progress_manager)
|
|
||||||
tracker = HFProgressTracker(progress_callback, filter_non_downloads=is_cached)
|
|
||||||
|
|
||||||
# Patch tqdm BEFORE importing mlx_audio
|
|
||||||
tracker_context = tracker.patch_download()
|
|
||||||
tracker_context.__enter__()
|
|
||||||
|
|
||||||
# Import mlx_audio
|
|
||||||
from mlx_audio.stt import load
|
from mlx_audio.stt import load
|
||||||
|
|
||||||
# MLX Whisper uses the standard OpenAI models
|
|
||||||
model_name = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}")
|
model_name = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}")
|
||||||
|
|
||||||
print(f"Loading MLX Whisper model {model_size}...")
|
print(f"Loading MLX Whisper model {model_size}...")
|
||||||
|
self.model = load(model_name)
|
||||||
# Only track download progress if model is NOT cached
|
|
||||||
if not is_cached:
|
self.model_size = model_size
|
||||||
# Start tracking download task
|
print(f"MLX Whisper model {model_size} loaded successfully")
|
||||||
task_manager.start_download(progress_model_name)
|
|
||||||
|
|
||||||
# Initialize progress state so SSE endpoint has initial data to send
|
|
||||||
progress_manager.update_progress(
|
|
||||||
model_name=progress_model_name,
|
|
||||||
current=0,
|
|
||||||
total=0,
|
|
||||||
filename="Connecting to HuggingFace...",
|
|
||||||
status="downloading",
|
|
||||||
)
|
|
||||||
|
|
||||||
# Load the model (tqdm is patched, but filters out non-download progress)
|
|
||||||
try:
|
|
||||||
self.model = load(model_name)
|
|
||||||
finally:
|
|
||||||
# Exit the patch context
|
|
||||||
tracker_context.__exit__(None, None, None)
|
|
||||||
|
|
||||||
# Only mark download as complete if we were tracking it
|
|
||||||
if not is_cached:
|
|
||||||
progress_manager.mark_complete(progress_model_name)
|
|
||||||
task_manager.complete_download(progress_model_name)
|
|
||||||
|
|
||||||
self.model_size = model_size
|
|
||||||
|
|
||||||
print(f"MLX Whisper model {model_size} loaded successfully")
|
|
||||||
|
|
||||||
except ImportError as e:
|
|
||||||
print(f"Error: mlx_audio package not found. Install with: pip install mlx-audio")
|
|
||||||
progress_manager = get_progress_manager()
|
|
||||||
task_manager = get_task_manager()
|
|
||||||
progress_model_name = f"whisper-{model_size}"
|
|
||||||
progress_manager.mark_error(progress_model_name, str(e))
|
|
||||||
task_manager.error_download(progress_model_name, str(e))
|
|
||||||
raise
|
|
||||||
except Exception as e:
|
|
||||||
print(f"Error loading MLX Whisper model: {e}")
|
|
||||||
progress_manager = get_progress_manager()
|
|
||||||
task_manager = get_task_manager()
|
|
||||||
progress_model_name = f"whisper-{model_size}"
|
|
||||||
progress_manager.mark_error(progress_model_name, str(e))
|
|
||||||
task_manager.error_download(progress_model_name, str(e))
|
|
||||||
raise
|
|
||||||
|
|
||||||
def unload_model(self):
|
def unload_model(self):
|
||||||
"""Unload the model to free memory."""
|
"""Unload the model to free memory."""
|
||||||
|
|||||||
@@ -6,20 +6,11 @@ from typing import Optional, List, Tuple
|
|||||||
import asyncio
|
import asyncio
|
||||||
import torch
|
import torch
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
from . import TTSBackend, STTBackend
|
from . import TTSBackend, STTBackend, LANGUAGE_CODE_TO_NAME, WHISPER_HF_REPOS
|
||||||
|
from .base import is_model_cached, get_torch_device, combine_voice_prompts as _combine_voice_prompts, model_load_progress
|
||||||
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
|
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
|
||||||
from ..utils.audio import normalize_audio, load_audio
|
from ..utils.audio import load_audio
|
||||||
from ..utils.progress import get_progress_manager
|
|
||||||
from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback
|
|
||||||
from ..utils.tasks import get_task_manager
|
|
||||||
|
|
||||||
LANGUAGE_CODE_TO_NAME = {
|
|
||||||
"zh": "chinese", "en": "english", "ja": "japanese", "ko": "korean",
|
|
||||||
"de": "german", "fr": "french", "ru": "russian", "pt": "portuguese",
|
|
||||||
"es": "spanish", "it": "italian",
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
class PyTorchTTSBackend:
|
class PyTorchTTSBackend:
|
||||||
@@ -33,26 +24,7 @@ class PyTorchTTSBackend:
|
|||||||
|
|
||||||
def _get_device(self) -> str:
|
def _get_device(self) -> str:
|
||||||
"""Get the best available device."""
|
"""Get the best available device."""
|
||||||
if torch.cuda.is_available():
|
return get_torch_device(allow_xpu=True, allow_directml=True)
|
||||||
return "cuda"
|
|
||||||
# Intel Arc / Intel Xe GPU via intel-extension-for-pytorch (IPEX)
|
|
||||||
try:
|
|
||||||
import intel_extension_for_pytorch # noqa: F401
|
|
||||||
if hasattr(torch, 'xpu') and torch.xpu.is_available():
|
|
||||||
return "xpu"
|
|
||||||
except ImportError:
|
|
||||||
pass
|
|
||||||
# Any GPU on Windows via DirectML (torch-directml)
|
|
||||||
try:
|
|
||||||
import torch_directml
|
|
||||||
if torch_directml.device_count() > 0:
|
|
||||||
return torch_directml.device(0)
|
|
||||||
except ImportError:
|
|
||||||
pass
|
|
||||||
# MPS (Apple Silicon) — kept for completeness but MLX backend is preferred
|
|
||||||
if hasattr(torch.backends, 'mps') and torch.backends.mps.is_available():
|
|
||||||
return "cpu" # MPS disabled for stability; MLX backend handles Apple Silicon
|
|
||||||
return "cpu"
|
|
||||||
|
|
||||||
def is_loaded(self) -> bool:
|
def is_loaded(self) -> bool:
|
||||||
"""Check if model is loaded."""
|
"""Check if model is loaded."""
|
||||||
@@ -79,44 +51,7 @@ class PyTorchTTSBackend:
|
|||||||
return hf_model_map[model_size]
|
return hf_model_map[model_size]
|
||||||
|
|
||||||
def _is_model_cached(self, model_size: str) -> bool:
|
def _is_model_cached(self, model_size: str) -> bool:
|
||||||
"""
|
return is_model_cached(self._get_model_path(model_size))
|
||||||
Check if the model is already cached locally AND fully downloaded.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
model_size: Model size to check
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True if model is fully cached, False if missing or incomplete
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
from huggingface_hub import constants as hf_constants
|
|
||||||
model_path = self._get_model_path(model_size)
|
|
||||||
repo_cache = Path(hf_constants.HF_HUB_CACHE) / ("models--" + model_path.replace("/", "--"))
|
|
||||||
|
|
||||||
if not repo_cache.exists():
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Check for .incomplete files - if any exist, download is still in progress
|
|
||||||
blobs_dir = repo_cache / "blobs"
|
|
||||||
if blobs_dir.exists() and any(blobs_dir.glob("*.incomplete")):
|
|
||||||
print(f"[_is_model_cached] Found .incomplete files for {model_size}, treating as not cached")
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Check that actual model weight files exist in snapshots
|
|
||||||
snapshots_dir = repo_cache / "snapshots"
|
|
||||||
if snapshots_dir.exists():
|
|
||||||
has_weights = (
|
|
||||||
any(snapshots_dir.rglob("*.safetensors")) or
|
|
||||||
any(snapshots_dir.rglob("*.bin"))
|
|
||||||
)
|
|
||||||
if not has_weights:
|
|
||||||
print(f"[_is_model_cached] No model weights found for {model_size}, treating as not cached")
|
|
||||||
return False
|
|
||||||
|
|
||||||
return True
|
|
||||||
except Exception as e:
|
|
||||||
print(f"[_is_model_cached] Error checking cache for {model_size}: {e}")
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def load_model_async(self, model_size: Optional[str] = None):
|
async def load_model_async(self, model_size: Optional[str] = None):
|
||||||
"""
|
"""
|
||||||
@@ -144,94 +79,30 @@ class PyTorchTTSBackend:
|
|||||||
|
|
||||||
def _load_model_sync(self, model_size: str):
|
def _load_model_sync(self, model_size: str):
|
||||||
"""Synchronous model loading."""
|
"""Synchronous model loading."""
|
||||||
try:
|
model_name = f"qwen-tts-{model_size}"
|
||||||
progress_manager = get_progress_manager()
|
is_cached = self._is_model_cached(model_size)
|
||||||
task_manager = get_task_manager()
|
|
||||||
model_name = f"qwen-tts-{model_size}"
|
|
||||||
|
|
||||||
# Check if model is already cached
|
with model_load_progress(model_name, is_cached):
|
||||||
is_cached = self._is_model_cached(model_size)
|
|
||||||
|
|
||||||
# Set up progress callback and tracker
|
|
||||||
# If cached: filter out non-download progress (like "Segment 1/1" during generation)
|
|
||||||
# If not cached: report all progress (we're actually downloading)
|
|
||||||
progress_callback = create_hf_progress_callback(model_name, progress_manager)
|
|
||||||
tracker = HFProgressTracker(progress_callback, filter_non_downloads=is_cached)
|
|
||||||
|
|
||||||
# Patch tqdm BEFORE importing qwen_tts
|
|
||||||
tracker_context = tracker.patch_download()
|
|
||||||
tracker_context.__enter__()
|
|
||||||
|
|
||||||
# Import qwen_tts
|
|
||||||
from qwen_tts import Qwen3TTSModel
|
from qwen_tts import Qwen3TTSModel
|
||||||
|
|
||||||
# Get model path (local or HuggingFace Hub ID)
|
|
||||||
model_path = self._get_model_path(model_size)
|
model_path = self._get_model_path(model_size)
|
||||||
|
|
||||||
print(f"Loading TTS model {model_size} on {self.device}...")
|
print(f"Loading TTS model {model_size} on {self.device}...")
|
||||||
|
|
||||||
# Only track download progress if model is NOT cached
|
if self.device == "cpu":
|
||||||
if not is_cached:
|
self.model = Qwen3TTSModel.from_pretrained(
|
||||||
# Start tracking download task
|
model_path,
|
||||||
task_manager.start_download(model_name)
|
torch_dtype=torch.float32,
|
||||||
|
low_cpu_mem_usage=False,
|
||||||
# Initialize progress state so SSE endpoint has initial data to send
|
)
|
||||||
progress_manager.update_progress(
|
else:
|
||||||
model_name=model_name,
|
self.model = Qwen3TTSModel.from_pretrained(
|
||||||
current=0,
|
model_path,
|
||||||
total=0, # Will be updated once actual total is known
|
device_map=self.device,
|
||||||
filename="Connecting to HuggingFace...",
|
torch_dtype=torch.bfloat16,
|
||||||
status="downloading",
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# Load the model (tqdm is patched, but filters out non-download progress)
|
self._current_model_size = model_size
|
||||||
try:
|
self.model_size = model_size
|
||||||
# Don't pass device_map on CPU: accelerate's meta-tensor mechanism
|
print(f"TTS model {model_size} loaded successfully")
|
||||||
# causes "Cannot copy out of meta tensor" when moving to CPU.
|
|
||||||
# Instead load directly then call .to(device) if needed.
|
|
||||||
if self.device == "cpu":
|
|
||||||
self.model = Qwen3TTSModel.from_pretrained(
|
|
||||||
model_path,
|
|
||||||
torch_dtype=torch.float32,
|
|
||||||
low_cpu_mem_usage=False,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
self.model = Qwen3TTSModel.from_pretrained(
|
|
||||||
model_path,
|
|
||||||
device_map=self.device,
|
|
||||||
torch_dtype=torch.bfloat16,
|
|
||||||
)
|
|
||||||
finally:
|
|
||||||
# Exit the patch context
|
|
||||||
tracker_context.__exit__(None, None, None)
|
|
||||||
|
|
||||||
# Only mark download as complete if we were tracking it
|
|
||||||
if not is_cached:
|
|
||||||
progress_manager.mark_complete(model_name)
|
|
||||||
task_manager.complete_download(model_name)
|
|
||||||
|
|
||||||
self._current_model_size = model_size
|
|
||||||
self.model_size = model_size
|
|
||||||
|
|
||||||
print(f"TTS model {model_size} loaded successfully")
|
|
||||||
|
|
||||||
except ImportError as e:
|
|
||||||
print(f"Error: qwen_tts package not found. Install with: pip install git+https://github.com/QwenLM/Qwen3-TTS.git")
|
|
||||||
progress_manager = get_progress_manager()
|
|
||||||
task_manager = get_task_manager()
|
|
||||||
model_name = f"qwen-tts-{model_size}"
|
|
||||||
progress_manager.mark_error(model_name, str(e))
|
|
||||||
task_manager.error_download(model_name, str(e))
|
|
||||||
raise
|
|
||||||
except Exception as e:
|
|
||||||
print(f"Error loading TTS model: {e}")
|
|
||||||
print(f"Tip: The model will be automatically downloaded from HuggingFace Hub on first use.")
|
|
||||||
progress_manager = get_progress_manager()
|
|
||||||
task_manager = get_task_manager()
|
|
||||||
model_name = f"qwen-tts-{model_size}"
|
|
||||||
progress_manager.mark_error(model_name, str(e))
|
|
||||||
task_manager.error_download(model_name, str(e))
|
|
||||||
raise
|
|
||||||
|
|
||||||
def unload_model(self):
|
def unload_model(self):
|
||||||
"""Unload the model to free memory."""
|
"""Unload the model to free memory."""
|
||||||
@@ -303,31 +174,7 @@ class PyTorchTTSBackend:
|
|||||||
audio_paths: List[str],
|
audio_paths: List[str],
|
||||||
reference_texts: List[str],
|
reference_texts: List[str],
|
||||||
) -> Tuple[np.ndarray, str]:
|
) -> Tuple[np.ndarray, str]:
|
||||||
"""
|
return await _combine_voice_prompts(audio_paths, reference_texts)
|
||||||
Combine multiple reference samples for better quality.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
audio_paths: List of audio file paths
|
|
||||||
reference_texts: List of reference texts
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tuple of (combined_audio, combined_text)
|
|
||||||
"""
|
|
||||||
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 generate(
|
async def generate(
|
||||||
self,
|
self,
|
||||||
@@ -376,15 +223,6 @@ class PyTorchTTSBackend:
|
|||||||
return audio, sample_rate
|
return audio, sample_rate
|
||||||
|
|
||||||
|
|
||||||
WHISPER_HF_REPOS = {
|
|
||||||
"base": "openai/whisper-base",
|
|
||||||
"small": "openai/whisper-small",
|
|
||||||
"medium": "openai/whisper-medium",
|
|
||||||
"large": "openai/whisper-large-v3",
|
|
||||||
"turbo": "openai/whisper-large-v3-turbo",
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
class PyTorchSTTBackend:
|
class PyTorchSTTBackend:
|
||||||
"""PyTorch-based STT backend using Whisper."""
|
"""PyTorch-based STT backend using Whisper."""
|
||||||
|
|
||||||
@@ -396,69 +234,15 @@ class PyTorchSTTBackend:
|
|||||||
|
|
||||||
def _get_device(self) -> str:
|
def _get_device(self) -> str:
|
||||||
"""Get the best available device."""
|
"""Get the best available device."""
|
||||||
if torch.cuda.is_available():
|
return get_torch_device(allow_xpu=True, allow_directml=True)
|
||||||
return "cuda"
|
|
||||||
# Intel Arc / Intel Xe GPU via intel-extension-for-pytorch (IPEX)
|
|
||||||
try:
|
|
||||||
import intel_extension_for_pytorch # noqa: F401
|
|
||||||
if hasattr(torch, 'xpu') and torch.xpu.is_available():
|
|
||||||
return "xpu"
|
|
||||||
except ImportError:
|
|
||||||
pass
|
|
||||||
# Any GPU on Windows via DirectML (torch-directml)
|
|
||||||
try:
|
|
||||||
import torch_directml
|
|
||||||
if torch_directml.device_count() > 0:
|
|
||||||
return torch_directml.device(0)
|
|
||||||
except ImportError:
|
|
||||||
pass
|
|
||||||
if hasattr(torch.backends, 'mps') and torch.backends.mps.is_available():
|
|
||||||
return "cpu" # MPS disabled for stability
|
|
||||||
return "cpu"
|
|
||||||
|
|
||||||
def is_loaded(self) -> bool:
|
def is_loaded(self) -> bool:
|
||||||
"""Check if model is loaded."""
|
"""Check if model is loaded."""
|
||||||
return self.model is not None
|
return self.model is not None
|
||||||
|
|
||||||
def _is_model_cached(self, model_size: str) -> bool:
|
def _is_model_cached(self, model_size: str) -> bool:
|
||||||
"""
|
hf_repo = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}")
|
||||||
Check if the Whisper model is already cached locally AND fully downloaded.
|
return is_model_cached(hf_repo)
|
||||||
|
|
||||||
Args:
|
|
||||||
model_size: Model size to check
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
True if model is fully cached, False if missing or incomplete
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
from huggingface_hub import constants as hf_constants
|
|
||||||
hf_repo = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}")
|
|
||||||
repo_cache = Path(hf_constants.HF_HUB_CACHE) / ("models--" + hf_repo.replace("/", "--"))
|
|
||||||
|
|
||||||
if not repo_cache.exists():
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Check for .incomplete files - if any exist, download is still in progress
|
|
||||||
blobs_dir = repo_cache / "blobs"
|
|
||||||
if blobs_dir.exists() and any(blobs_dir.glob("*.incomplete")):
|
|
||||||
print(f"[_is_model_cached] Found .incomplete files for whisper-{model_size}, treating as not cached")
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Check that actual model weight files exist in snapshots
|
|
||||||
snapshots_dir = repo_cache / "snapshots"
|
|
||||||
if snapshots_dir.exists():
|
|
||||||
has_weights = (
|
|
||||||
any(snapshots_dir.rglob("*.safetensors")) or
|
|
||||||
any(snapshots_dir.rglob("*.bin"))
|
|
||||||
)
|
|
||||||
if not has_weights:
|
|
||||||
print(f"[_is_model_cached] No model weights found for whisper-{model_size}, treating as not cached")
|
|
||||||
return False
|
|
||||||
|
|
||||||
return True
|
|
||||||
except Exception as e:
|
|
||||||
print(f"[_is_model_cached] Error checking cache for whisper-{model_size}: {e}")
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def load_model_async(self, model_size: Optional[str] = None):
|
async def load_model_async(self, model_size: Optional[str] = None):
|
||||||
"""
|
"""
|
||||||
@@ -467,94 +251,33 @@ class PyTorchSTTBackend:
|
|||||||
Args:
|
Args:
|
||||||
model_size: Model size (tiny, base, small, medium, large)
|
model_size: Model size (tiny, base, small, medium, large)
|
||||||
"""
|
"""
|
||||||
print(f"[DEBUG] load_model_async called with size: {model_size}")
|
|
||||||
if model_size is None:
|
if model_size is None:
|
||||||
model_size = self.model_size
|
model_size = self.model_size
|
||||||
|
|
||||||
print(f"[DEBUG] Model already loaded? {self.model is not None}, current size: {self.model_size}, requested: {model_size}")
|
|
||||||
if self.model is not None and self.model_size == model_size:
|
if self.model is not None and self.model_size == model_size:
|
||||||
print(f"[DEBUG] Early return - model already loaded")
|
|
||||||
return
|
return
|
||||||
|
|
||||||
print(f"[DEBUG] Calling asyncio.to_thread for _load_model_sync")
|
|
||||||
# Run blocking load in thread pool
|
|
||||||
await asyncio.to_thread(self._load_model_sync, model_size)
|
await asyncio.to_thread(self._load_model_sync, model_size)
|
||||||
print(f"[DEBUG] asyncio.to_thread completed")
|
|
||||||
|
|
||||||
# Alias for compatibility
|
# Alias for compatibility
|
||||||
load_model = load_model_async
|
load_model = load_model_async
|
||||||
|
|
||||||
def _load_model_sync(self, model_size: str):
|
def _load_model_sync(self, model_size: str):
|
||||||
"""Synchronous model loading."""
|
"""Synchronous model loading."""
|
||||||
print(f"[DEBUG] _load_model_sync called for Whisper {model_size}")
|
progress_model_name = f"whisper-{model_size}"
|
||||||
try:
|
is_cached = self._is_model_cached(model_size)
|
||||||
progress_manager = get_progress_manager()
|
|
||||||
task_manager = get_task_manager()
|
|
||||||
progress_model_name = f"whisper-{model_size}"
|
|
||||||
|
|
||||||
# Check if model is already cached
|
with model_load_progress(progress_model_name, is_cached):
|
||||||
is_cached = self._is_model_cached(model_size)
|
|
||||||
|
|
||||||
# Set up progress callback and tracker
|
|
||||||
# If cached: filter out non-download progress
|
|
||||||
# If not cached: report all progress (we're actually downloading)
|
|
||||||
progress_callback = create_hf_progress_callback(progress_model_name, progress_manager)
|
|
||||||
tracker = HFProgressTracker(progress_callback, filter_non_downloads=is_cached)
|
|
||||||
|
|
||||||
# Patch tqdm BEFORE importing transformers
|
|
||||||
print("[DEBUG] Starting tqdm patch BEFORE transformers import")
|
|
||||||
tracker_context = tracker.patch_download()
|
|
||||||
tracker_context.__enter__()
|
|
||||||
print("[DEBUG] tqdm patched, now importing transformers")
|
|
||||||
|
|
||||||
# Import transformers
|
|
||||||
from transformers import WhisperProcessor, WhisperForConditionalGeneration
|
from transformers import WhisperProcessor, WhisperForConditionalGeneration
|
||||||
|
|
||||||
model_name = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}")
|
model_name = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}")
|
||||||
print(f"[DEBUG] Model name: {model_name}")
|
|
||||||
|
|
||||||
print(f"Loading Whisper model {model_size} on {self.device}...")
|
print(f"Loading Whisper model {model_size} on {self.device}...")
|
||||||
|
|
||||||
# Only track download progress if model is NOT cached
|
self.processor = WhisperProcessor.from_pretrained(model_name)
|
||||||
if not is_cached:
|
self.model = WhisperForConditionalGeneration.from_pretrained(model_name)
|
||||||
# Start tracking download task
|
|
||||||
task_manager.start_download(progress_model_name)
|
|
||||||
|
|
||||||
# Initialize progress state so SSE endpoint has initial data to send
|
self.model.to(self.device)
|
||||||
progress_manager.update_progress(
|
self.model_size = model_size
|
||||||
model_name=progress_model_name,
|
print(f"Whisper model {model_size} loaded successfully")
|
||||||
current=0,
|
|
||||||
total=0, # Will be updated once actual total is known
|
|
||||||
filename="Connecting to HuggingFace...",
|
|
||||||
status="downloading",
|
|
||||||
)
|
|
||||||
|
|
||||||
# Load models (tqdm is patched, but filters out non-download progress)
|
|
||||||
try:
|
|
||||||
self.processor = WhisperProcessor.from_pretrained(model_name)
|
|
||||||
self.model = WhisperForConditionalGeneration.from_pretrained(model_name)
|
|
||||||
finally:
|
|
||||||
# Exit the patch context
|
|
||||||
tracker_context.__exit__(None, None, None)
|
|
||||||
|
|
||||||
# Only mark download as complete if we were tracking it
|
|
||||||
if not is_cached:
|
|
||||||
progress_manager.mark_complete(progress_model_name)
|
|
||||||
task_manager.complete_download(progress_model_name)
|
|
||||||
|
|
||||||
self.model.to(self.device)
|
|
||||||
self.model_size = model_size
|
|
||||||
|
|
||||||
print(f"Whisper model {model_size} loaded successfully")
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
print(f"Error loading Whisper model: {e}")
|
|
||||||
progress_manager = get_progress_manager()
|
|
||||||
task_manager = get_task_manager()
|
|
||||||
progress_model_name = f"whisper-{model_size}"
|
|
||||||
progress_manager.mark_error(progress_model_name, str(e))
|
|
||||||
task_manager.error_download(progress_model_name, str(e))
|
|
||||||
raise
|
|
||||||
|
|
||||||
def unload_model(self):
|
def unload_model(self):
|
||||||
"""Unload the model to free memory."""
|
"""Unload the model to free memory."""
|
||||||
|
|||||||
@@ -19,11 +19,6 @@ from .models import VoiceProfileCreate
|
|||||||
from . import config
|
from . import config
|
||||||
|
|
||||||
|
|
||||||
def _get_profiles_dir() -> Path:
|
|
||||||
"""Get profiles directory from config."""
|
|
||||||
return config.get_profiles_dir()
|
|
||||||
|
|
||||||
|
|
||||||
def _get_unique_profile_name(name: str, db: Session) -> str:
|
def _get_unique_profile_name(name: str, db: Session) -> str:
|
||||||
"""
|
"""
|
||||||
Get a unique profile name by appending a number if needed.
|
Get a unique profile name by appending a number if needed.
|
||||||
@@ -99,7 +94,7 @@ def export_profile_to_zip(profile_id: str, db: Session) -> bytes:
|
|||||||
|
|
||||||
# Create samples.json mapping
|
# Create samples.json mapping
|
||||||
samples_data = {}
|
samples_data = {}
|
||||||
profile_dir = _get_profiles_dir() / profile_id
|
profile_dir = config.get_profiles_dir() / profile_id
|
||||||
|
|
||||||
for sample in samples:
|
for sample in samples:
|
||||||
# Get filename from audio_path (should be {sample_id}.wav)
|
# Get filename from audio_path (should be {sample_id}.wav)
|
||||||
@@ -181,7 +176,7 @@ async def import_profile_from_zip(file_bytes: bytes, db: Session) -> VoiceProfil
|
|||||||
profile = await create_profile(profile_create, db)
|
profile = await create_profile(profile_create, db)
|
||||||
|
|
||||||
# Extract and add samples
|
# Extract and add samples
|
||||||
profile_dir = _get_profiles_dir() / profile.id
|
profile_dir = config.get_profiles_dir() / profile.id
|
||||||
profile_dir.mkdir(parents=True, exist_ok=True)
|
profile_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
# Handle avatar if present
|
# Handle avatar if present
|
||||||
|
|||||||
@@ -15,11 +15,6 @@ from .database import Generation as DBGeneration, GenerationVersion as DBGenerat
|
|||||||
from . import config
|
from . import config
|
||||||
|
|
||||||
|
|
||||||
def _get_generations_dir() -> Path:
|
|
||||||
"""Get generations directory from config."""
|
|
||||||
return config.get_generations_dir()
|
|
||||||
|
|
||||||
|
|
||||||
def _get_versions_for_generation(generation_id: str, db: Session) -> tuple:
|
def _get_versions_for_generation(generation_id: str, db: Session) -> tuple:
|
||||||
"""Get versions list and active version ID for a generation."""
|
"""Get versions list and active version ID for a generation."""
|
||||||
import json
|
import json
|
||||||
|
|||||||
+1
-26
@@ -18,7 +18,6 @@ import tempfile
|
|||||||
import io
|
import io
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
import uuid
|
import uuid
|
||||||
import asyncio
|
|
||||||
import signal
|
import signal
|
||||||
import os
|
import os
|
||||||
|
|
||||||
@@ -54,6 +53,7 @@ def _safe_content_disposition(disposition_type: str, filename: str) -> str:
|
|||||||
|
|
||||||
from . import database, models, profiles, history, tts, transcribe, config, export_import, channels, stories, __version__
|
from . import database, models, profiles, history, tts, transcribe, config, export_import, channels, stories, __version__
|
||||||
from .database import get_db, Generation as DBGeneration, VoiceProfile as DBVoiceProfile
|
from .database import get_db, Generation as DBGeneration, VoiceProfile as DBVoiceProfile
|
||||||
|
from .profiles import _profile_to_response
|
||||||
from .utils.progress import get_progress_manager
|
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
|
||||||
@@ -1980,31 +1980,6 @@ async def update_profile_effects(
|
|||||||
return _profile_to_response(profile)
|
return _profile_to_response(profile)
|
||||||
|
|
||||||
|
|
||||||
def _profile_to_response(profile) -> models.VoiceProfileResponse:
|
|
||||||
"""Convert a DB profile to a VoiceProfileResponse with parsed effects_chain."""
|
|
||||||
import json as _json
|
|
||||||
import logging
|
|
||||||
|
|
||||||
effects_chain = None
|
|
||||||
if profile.effects_chain:
|
|
||||||
try:
|
|
||||||
raw = _json.loads(profile.effects_chain)
|
|
||||||
effects_chain = [models.EffectConfig(**e) for e in raw]
|
|
||||||
except Exception as e:
|
|
||||||
logging.warning(f"Failed to parse effects_chain for profile {profile.id}: {e}")
|
|
||||||
|
|
||||||
return models.VoiceProfileResponse(
|
|
||||||
id=profile.id,
|
|
||||||
name=profile.name,
|
|
||||||
description=profile.description,
|
|
||||||
language=profile.language,
|
|
||||||
avatar_path=profile.avatar_path,
|
|
||||||
effects_chain=effects_chain,
|
|
||||||
created_at=profile.created_at,
|
|
||||||
updated_at=profile.updated_at,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
# ============================================
|
# ============================================
|
||||||
# FILE SERVING
|
# FILE SERVING
|
||||||
# ============================================
|
# ============================================
|
||||||
|
|||||||
@@ -1,48 +0,0 @@
|
|||||||
"""
|
|
||||||
Database migration script to add instruct column to generations table.
|
|
||||||
|
|
||||||
Run this once to update existing databases:
|
|
||||||
python -m backend.migrate_add_instruct
|
|
||||||
"""
|
|
||||||
|
|
||||||
import sqlite3
|
|
||||||
import os
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
|
|
||||||
def migrate():
|
|
||||||
"""Add instruct column to generations table if it doesn't exist."""
|
|
||||||
# Get data directory
|
|
||||||
data_dir = os.environ.get("VOICEBOX_DATA_DIR")
|
|
||||||
if data_dir:
|
|
||||||
db_path = Path(data_dir) / "voicebox.db"
|
|
||||||
else:
|
|
||||||
db_path = Path.cwd() / "data" / "voicebox.db"
|
|
||||||
|
|
||||||
if not db_path.exists():
|
|
||||||
print(f"Database not found at {db_path}, skipping migration")
|
|
||||||
return
|
|
||||||
|
|
||||||
conn = sqlite3.connect(db_path)
|
|
||||||
cursor = conn.cursor()
|
|
||||||
|
|
||||||
# Check if instruct column already exists
|
|
||||||
cursor.execute("PRAGMA table_info(generations)")
|
|
||||||
columns = [row[1] for row in cursor.fetchall()]
|
|
||||||
|
|
||||||
if 'instruct' in columns:
|
|
||||||
print("instruct column already exists, skipping migration")
|
|
||||||
conn.close()
|
|
||||||
return
|
|
||||||
|
|
||||||
# Add instruct column
|
|
||||||
print("Adding instruct column to generations table...")
|
|
||||||
cursor.execute("ALTER TABLE generations ADD COLUMN instruct TEXT")
|
|
||||||
conn.commit()
|
|
||||||
conn.close()
|
|
||||||
|
|
||||||
print("Migration complete!")
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
migrate()
|
|
||||||
+4
-9
@@ -58,11 +58,6 @@ def _profile_to_response(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _get_profiles_dir() -> Path:
|
|
||||||
"""Get profiles directory from config."""
|
|
||||||
return config.get_profiles_dir()
|
|
||||||
|
|
||||||
|
|
||||||
async def create_profile(
|
async def create_profile(
|
||||||
data: VoiceProfileCreate,
|
data: VoiceProfileCreate,
|
||||||
db: Session,
|
db: Session,
|
||||||
@@ -100,7 +95,7 @@ async def create_profile(
|
|||||||
db.refresh(db_profile)
|
db.refresh(db_profile)
|
||||||
|
|
||||||
# Create profile directory
|
# Create profile directory
|
||||||
profile_dir = _get_profiles_dir() / db_profile.id
|
profile_dir = config.get_profiles_dir() / db_profile.id
|
||||||
profile_dir.mkdir(parents=True, exist_ok=True)
|
profile_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
return _profile_to_response(db_profile)
|
return _profile_to_response(db_profile)
|
||||||
@@ -136,7 +131,7 @@ async def add_profile_sample(
|
|||||||
|
|
||||||
# Create sample ID and directory
|
# Create sample ID and directory
|
||||||
sample_id = str(uuid.uuid4())
|
sample_id = str(uuid.uuid4())
|
||||||
profile_dir = _get_profiles_dir() / profile_id
|
profile_dir = config.get_profiles_dir() / profile_id
|
||||||
profile_dir.mkdir(parents=True, exist_ok=True)
|
profile_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
# Copy audio file to profile directory
|
# Copy audio file to profile directory
|
||||||
@@ -316,7 +311,7 @@ async def delete_profile(
|
|||||||
db.commit()
|
db.commit()
|
||||||
|
|
||||||
# Delete profile directory
|
# Delete profile directory
|
||||||
profile_dir = _get_profiles_dir() / profile_id
|
profile_dir = config.get_profiles_dir() / profile_id
|
||||||
if profile_dir.exists():
|
if profile_dir.exists():
|
||||||
shutil.rmtree(profile_dir)
|
shutil.rmtree(profile_dir)
|
||||||
|
|
||||||
@@ -516,7 +511,7 @@ async def upload_avatar(
|
|||||||
ext = ext_map.get(img_format, '.png')
|
ext = ext_map.get(img_format, '.png')
|
||||||
|
|
||||||
# Save processed image to profile directory
|
# Save processed image to profile directory
|
||||||
profile_dir = _get_profiles_dir() / profile_id
|
profile_dir = config.get_profiles_dir() / profile_id
|
||||||
profile_dir.mkdir(parents=True, exist_ok=True)
|
profile_dir.mkdir(parents=True, exist_ok=True)
|
||||||
output_path = profile_dir / f"avatar{ext}"
|
output_path = profile_dir / f"avatar{ext}"
|
||||||
|
|
||||||
|
|||||||
@@ -1,66 +0,0 @@
|
|||||||
"""
|
|
||||||
Audio studio module for timeline editing.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from typing import List, Dict, Optional
|
|
||||||
import numpy as np
|
|
||||||
|
|
||||||
|
|
||||||
class AudioStudio:
|
|
||||||
"""Audio editing and timeline management."""
|
|
||||||
|
|
||||||
async def get_word_timestamps(
|
|
||||||
self,
|
|
||||||
audio_path: str,
|
|
||||||
text: str,
|
|
||||||
) -> List[Dict[str, float]]:
|
|
||||||
"""
|
|
||||||
Get word-level timestamps for audio.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
audio_path: Path to audio file
|
|
||||||
text: Corresponding text
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
List of word timestamps: [{"word": "...", "start": 0.0, "end": 0.5}, ...]
|
|
||||||
"""
|
|
||||||
# TODO: Implement Whisper alignment
|
|
||||||
raise NotImplementedError("Word timestamps not yet implemented")
|
|
||||||
|
|
||||||
async def mix_audio(
|
|
||||||
self,
|
|
||||||
audio_paths: List[str],
|
|
||||||
volumes: Optional[List[float]] = None,
|
|
||||||
) -> bytes:
|
|
||||||
"""
|
|
||||||
Mix multiple audio files together.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
audio_paths: List of audio file paths
|
|
||||||
volumes: Optional volume levels (0.0-1.0) for each track
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Mixed audio bytes (WAV format)
|
|
||||||
"""
|
|
||||||
# TODO: Implement audio mixing
|
|
||||||
raise NotImplementedError("Audio mixing not yet implemented")
|
|
||||||
|
|
||||||
async def trim_audio(
|
|
||||||
self,
|
|
||||||
audio_path: str,
|
|
||||||
start: float,
|
|
||||||
end: float,
|
|
||||||
) -> bytes:
|
|
||||||
"""
|
|
||||||
Trim audio to specified time range.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
audio_path: Path to audio file
|
|
||||||
start: Start time in seconds
|
|
||||||
end: End time in seconds
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Trimmed audio bytes (WAV format)
|
|
||||||
"""
|
|
||||||
# TODO: Implement audio trimming
|
|
||||||
raise NotImplementedError("Audio trimming not yet implemented")
|
|
||||||
@@ -45,8 +45,8 @@ def test_db():
|
|||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def mock_profiles_dir(monkeypatch, tmp_path):
|
def mock_profiles_dir(monkeypatch, tmp_path):
|
||||||
"""Mock the profiles directory to use a temporary path."""
|
"""Mock the profiles directory to use a temporary path."""
|
||||||
import profiles
|
from backend import config
|
||||||
monkeypatch.setattr(profiles, '_get_profiles_dir', lambda: tmp_path)
|
monkeypatch.setattr(config, 'get_profiles_dir', lambda: tmp_path)
|
||||||
return tmp_path
|
return tmp_path
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,66 +0,0 @@
|
|||||||
"""
|
|
||||||
Input validation utilities.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from typing import Tuple, Optional
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
|
|
||||||
def validate_text(text: str, max_length: int = 5000) -> Tuple[bool, Optional[str]]:
|
|
||||||
"""
|
|
||||||
Validate text input.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
text: Text to validate
|
|
||||||
max_length: Maximum length
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tuple of (is_valid, error_message)
|
|
||||||
"""
|
|
||||||
if not text or not text.strip():
|
|
||||||
return False, "Text cannot be empty"
|
|
||||||
|
|
||||||
if len(text) > max_length:
|
|
||||||
return False, f"Text too long (maximum {max_length} characters)"
|
|
||||||
|
|
||||||
return True, None
|
|
||||||
|
|
||||||
|
|
||||||
def validate_language(language: str) -> Tuple[bool, Optional[str]]:
|
|
||||||
"""
|
|
||||||
Validate language code.
|
|
||||||
|
|
||||||
Supported languages for Qwen3-TTS:
|
|
||||||
Chinese, English, Japanese, Korean, German, French, Russian, Portuguese, Spanish, Italian
|
|
||||||
|
|
||||||
Args:
|
|
||||||
language: Language code
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tuple of (is_valid, error_message)
|
|
||||||
"""
|
|
||||||
valid_languages = ["zh", "en", "ja", "ko", "de", "fr", "ru", "pt", "es", "it"]
|
|
||||||
if language not in valid_languages:
|
|
||||||
return False, f"Invalid language (must be one of: {', '.join(valid_languages)})"
|
|
||||||
|
|
||||||
return True, None
|
|
||||||
|
|
||||||
|
|
||||||
def validate_file_path(path: str) -> Tuple[bool, Optional[str]]:
|
|
||||||
"""
|
|
||||||
Validate file path exists.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
path: File path
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Tuple of (is_valid, error_message)
|
|
||||||
"""
|
|
||||||
file_path = Path(path)
|
|
||||||
if not file_path.exists():
|
|
||||||
return False, f"File not found: {path}"
|
|
||||||
|
|
||||||
if not file_path.is_file():
|
|
||||||
return False, f"Path is not a file: {path}"
|
|
||||||
|
|
||||||
return True, None
|
|
||||||
Reference in New Issue
Block a user