From 0813a3d9d62e83eee686082262de85d084df6ac6 Mon Sep 17 00:00:00 2001 From: Jamie Pine Date: Mon, 16 Mar 2026 01:10:02 -0700 Subject: [PATCH] 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 --- backend/README.md | 51 ++- backend/backends/__init__.py | 14 + backend/backends/base.py | 277 ++++++++++++++ backend/backends/chatterbox_backend.py | 187 ++-------- backend/backends/chatterbox_turbo_backend.py | 176 +-------- backend/backends/luxtts_backend.py | 136 +------ backend/backends/mlx_backend.py | 337 +++-------------- backend/backends/pytorch_backend.py | 345 ++---------------- backend/export_import.py | 9 +- backend/history.py | 5 - backend/main.py | 27 +- backend/migrate_add_instruct.py | 48 --- backend/profiles.py | 13 +- backend/studio.py | 66 ---- backend/tests/test_profile_duplicate_names.py | 4 +- backend/utils/validation.py | 66 ---- 16 files changed, 480 insertions(+), 1281 deletions(-) create mode 100644 backend/backends/base.py delete mode 100644 backend/migrate_add_instruct.py delete mode 100644 backend/studio.py delete mode 100644 backend/utils/validation.py diff --git a/backend/README.md b/backend/README.md index 88f872b1..bfbb4ac1 100644 --- a/backend/README.md +++ b/backend/README.md @@ -17,23 +17,38 @@ Production-quality FastAPI backend for Qwen3-TTS voice cloning. ``` backend/ -├── main.py # FastAPI app with all routes -├── models.py # Pydantic request/response models -├── platform_detect.py # Platform detection for backend selection -├── tts.py # TTS backend abstraction (delegates to MLX or PyTorch) -├── transcribe.py # STT backend abstraction (delegates to MLX or PyTorch) -├── backends/ # Backend implementations -│ ├── __init__.py # Backend factory and protocols -│ ├── mlx_backend.py # MLX backend (Apple Silicon) -│ └── pytorch_backend.py # PyTorch backend (Windows/Linux/Intel) -├── profiles.py # Voice profile CRUD -├── history.py # Generation history -├── studio.py # Audio editing (TODO) -├── database.py # SQLite ORM +├── main.py # FastAPI app with all routes +├── server.py # PyInstaller entry point, CLI arg parsing +├── models.py # Pydantic request/response models +├── config.py # Data directory configuration +├── database.py # SQLAlchemy ORM models + migrations +├── platform_detect.py # Platform detection for backend selection +├── tts.py # TTS backend facade +├── transcribe.py # STT backend facade +├── profiles.py # Voice profile CRUD +├── history.py # Generation history CRUD +├── channels.py # Audio channel management +├── stories.py # Story/timeline management + audio export +├── 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/ - ├── audio.py # Audio processing utilities - ├── cache.py # Voice prompt caching - └── validation.py # Input validation + ├── audio.py # Audio load/save/normalize/validate/trim + ├── cache.py # Voice prompt caching (memory + disk) + ├── 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 @@ -450,12 +465,8 @@ Error responses include details: - [ ] WebSocket support for generation progress - [ ] Batch generation endpoint -- [ ] Audio effects (M3GAN, etc.) - [ ] Voice design (text-to-voice) -- [ ] Audio studio timeline features -- [ ] Project management - [ ] Authentication & rate limiting -- [ ] Export/import profiles ## License diff --git a/backend/backends/__init__.py b/backend/backends/__init__.py index e51cd0b3..c10f2241 100644 --- a/backend/backends/__init__.py +++ b/backend/backends/__init__.py @@ -13,6 +13,20 @@ import numpy as np 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 class ModelConfig: diff --git a/backend/backends/base.py b/backend/backends/base.py new file mode 100644 index 00000000..f839df35 --- /dev/null +++ b/backend/backends/base.py @@ -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) diff --git a/backend/backends/chatterbox_backend.py b/backend/backends/chatterbox_backend.py index bd6170b7..7efe371a 100644 --- a/backend/backends/chatterbox_backend.py +++ b/backend/backends/chatterbox_backend.py @@ -8,7 +8,6 @@ on macOS due to known MPS tensor issues. import asyncio import logging -import platform import threading from pathlib import Path from typing import ClassVar, List, Optional, Tuple @@ -16,9 +15,13 @@ from typing import ClassVar, List, Optional, Tuple import numpy as np from . import TTSBackend -from ..utils.audio import normalize_audio, load_audio -from ..utils.progress import get_progress_manager -from ..utils.tasks import get_task_manager +from .base import ( + is_model_cached, + get_torch_device, + combine_voice_prompts as _combine_voice_prompts, + model_load_progress, + patch_chatterbox_f32, +) logger = logging.getLogger(__name__) @@ -45,17 +48,7 @@ class ChatterboxTTSBackend: self._model_load_lock = asyncio.Lock() def _get_device(self) -> str: - """Get the best available device. Forces CPU on macOS (MPS issue).""" - if platform.system() == "Darwin": - return "cpu" - try: - import torch - - if torch.cuda.is_available(): - return "cuda" - except ImportError: - pass - return "cpu" + return get_torch_device(force_cpu_on_mac=True) def is_loaded(self) -> bool: return self.model is not None @@ -64,33 +57,7 @@ class ChatterboxTTSBackend: return CHATTERBOX_HF_REPO def _is_model_cached(self, model_size: str = "default") -> bool: - """Check if the Chatterbox multilingual model is cached locally.""" - 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 + return is_model_cached(CHATTERBOX_HF_REPO, required_files=_MTL_WEIGHT_FILES) async def load_model(self, model_size: str = "default") -> None: """Load the Chatterbox multilingual model.""" @@ -103,133 +70,45 @@ class ChatterboxTTSBackend: 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 = "chatterbox-tts" - is_cached = self._is_model_cached() - # Set up HF progress tracking (intercepts tqdm for file-level progress) - 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: + with model_load_progress(model_name, is_cached): device = self._get_device() self._device = device - logger.info(f"Loading Chatterbox Multilingual TTS on {device}...") import torch from chatterbox.mtl_tts import ChatterboxMultilingualTTS - # Load into a local variable first, apply all patches, then - # assign to self.model. This avoids leaving a half-initialised - # 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 + if device == "cpu": + _orig_torch_load = torch.load - def _patched_load(*args, **kwargs): - kwargs.setdefault("map_location", "cpu") - return _orig_torch_load(*args, **kwargs) + def _patched_load(*args, **kwargs): + kwargs.setdefault("map_location", "cpu") + return _orig_torch_load(*args, **kwargs) - with ChatterboxTTSBackend._load_lock: - torch.load = _patched_load - try: - model = ChatterboxMultilingualTTS.from_pretrained( - device=device, - ) - finally: - torch.load = _orig_torch_load - else: - model = ChatterboxMultilingualTTS.from_pretrained( - device=device, - ) - finally: - tracker_context.__exit__(None, None, None) + with ChatterboxTTSBackend._load_lock: + torch.load = _patched_load + try: + model = ChatterboxMultilingualTTS.from_pretrained(device=device) + finally: + torch.load = _orig_torch_load + else: + model = ChatterboxMultilingualTTS.from_pretrained(device=device) - # Fix: transformers >= 4.36 defaults LlamaModel to sdpa attention - # which doesn't support output_attentions=True (needed by - # Chatterbox's AlignmentStreamAnalyzer). Force eager attention. + # Fix sdpa attention for output_attentions support t3_tfmr = model.t3.tfmr - if hasattr(t3_tfmr, "config") and hasattr( - t3_tfmr.config, "_attn_implementation" - ): + if hasattr(t3_tfmr, "config") and hasattr(t3_tfmr.config, "_attn_implementation"): t3_tfmr.config._attn_implementation = "eager" for layer in getattr(t3_tfmr, "layers", []): if hasattr(layer, "self_attn"): layer.self_attn._attn_implementation = "eager" - if not is_cached: - 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 + patch_chatterbox_f32(model) self.model = model - 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 + logger.info("Chatterbox Multilingual TTS loaded successfully") def unload_model(self) -> None: """Unload model to free memory.""" @@ -268,17 +147,7 @@ class ChatterboxTTSBackend: audio_paths: List[str], reference_texts: List[str], ) -> Tuple[np.ndarray, str]: - """Combine multiple reference samples.""" - 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 + return await _combine_voice_prompts(audio_paths, reference_texts) # Per-language generation defaults. Lower temp + higher cfg = clearer speech. _LANG_DEFAULTS: ClassVar[dict] = { diff --git a/backend/backends/chatterbox_turbo_backend.py b/backend/backends/chatterbox_turbo_backend.py index 4971b8e9..a8bfe503 100644 --- a/backend/backends/chatterbox_turbo_backend.py +++ b/backend/backends/chatterbox_turbo_backend.py @@ -8,7 +8,6 @@ Forces CPU on macOS due to known MPS tensor issues. import asyncio import logging -import platform import threading from pathlib import Path from typing import ClassVar, List, Optional, Tuple @@ -16,9 +15,13 @@ from typing import ClassVar, List, Optional, Tuple import numpy as np from . import TTSBackend -from ..utils.audio import normalize_audio, load_audio -from ..utils.progress import get_progress_manager -from ..utils.tasks import get_task_manager +from .base import ( + is_model_cached, + get_torch_device, + combine_voice_prompts as _combine_voice_prompts, + model_load_progress, + patch_chatterbox_f32, +) logger = logging.getLogger(__name__) @@ -45,17 +48,7 @@ class ChatterboxTurboTTSBackend: self._model_load_lock = asyncio.Lock() def _get_device(self) -> str: - """Get the best available device. Forces CPU on macOS (MPS issue).""" - if platform.system() == "Darwin": - return "cpu" - try: - import torch - - if torch.cuda.is_available(): - return "cuda" - except ImportError: - pass - return "cpu" + return get_torch_device(force_cpu_on_mac=True) def is_loaded(self) -> bool: return self.model is not None @@ -64,33 +57,7 @@ class ChatterboxTurboTTSBackend: return CHATTERBOX_TURBO_HF_REPO def _is_model_cached(self, model_size: str = "default") -> bool: - """Check if the Chatterbox Turbo model is cached locally.""" - 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 + return is_model_cached(CHATTERBOX_TURBO_HF_REPO, required_files=_TURBO_WEIGHT_FILES) async def load_model(self, model_size: str = "default") -> None: """Load the Chatterbox Turbo model.""" @@ -103,59 +70,24 @@ class ChatterboxTurboTTSBackend: 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 = "chatterbox-turbo" - is_cached = self._is_model_cached() - # Set up HF progress tracking (intercepts tqdm for file-level progress) - 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: + with model_load_progress(model_name, is_cached): device = self._get_device() self._device = device - logger.info(f"Loading Chatterbox Turbo TTS on {device}...") import torch from huggingface_hub import snapshot_download from chatterbox.tts_turbo import ChatterboxTurboTTS - # Download model files ourselves so we can pass token=None - # (upstream from_pretrained passes token=True which requires - # a stored HF token even though the repo is public). - try: - 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) + local_path = snapshot_download( + repo_id=CHATTERBOX_TURBO_HF_REPO, + token=None, + allow_patterns=["*.safetensors", "*.json", "*.txt", "*.pt", "*.model"], + ) - # 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": _orig_torch_load = torch.load @@ -166,74 +98,16 @@ class ChatterboxTurboTTSBackend: with ChatterboxTurboTTSBackend._load_lock: torch.load = _patched_load try: - model = ChatterboxTurboTTS.from_local( - local_path, device, - ) + model = ChatterboxTurboTTS.from_local(local_path, device) finally: torch.load = _orig_torch_load else: - model = ChatterboxTurboTTS.from_local( - local_path, device, - ) + model = ChatterboxTurboTTS.from_local(local_path, device) - if not is_cached: - 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 + patch_chatterbox_f32(model) self.model = model - 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 + logger.info("Chatterbox Turbo TTS loaded successfully") def unload_model(self) -> None: """Unload model to free memory.""" @@ -271,17 +145,7 @@ class ChatterboxTurboTTSBackend: audio_paths: List[str], reference_texts: List[str], ) -> Tuple[np.ndarray, str]: - """Combine multiple reference samples.""" - 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 + return await _combine_voice_prompts(audio_paths, reference_texts) async def generate( self, diff --git a/backend/backends/luxtts_backend.py b/backend/backends/luxtts_backend.py index c289904f..ba00359e 100644 --- a/backend/backends/luxtts_backend.py +++ b/backend/backends/luxtts_backend.py @@ -7,16 +7,13 @@ Wraps the LuxTTS (ZipVoice) model for zero-shot voice cloning. import asyncio import logging -from pathlib import Path -from typing import List, Optional, Tuple +from typing import Optional, Tuple import numpy as np 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.progress import get_progress_manager -from ..utils.tasks import get_task_manager logger = logging.getLogger(__name__) @@ -33,14 +30,7 @@ class LuxTTSBackend: self._device = None def _get_device(self) -> str: - """Get the best available device.""" - import torch - - if torch.cuda.is_available(): - return "cuda" - if hasattr(torch.backends, "mps") and torch.backends.mps.is_available(): - return "mps" - return "cpu" + return get_torch_device(allow_mps=True) def is_loaded(self) -> bool: return self.model is not None @@ -55,35 +45,10 @@ class LuxTTSBackend: return LUXTTS_HF_REPO def _is_model_cached(self, model_size: str = "default") -> bool: - """Check if LuxTTS model weights are cached locally.""" - try: - from huggingface_hub import constants as hf_constants - - 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 + return is_model_cached( + LUXTTS_HF_REPO, + weight_extensions=(".pt", ".safetensors", ".onnx", ".bin"), + ) async def load_model(self, model_size: str = "default") -> None: """Load the LuxTTS model.""" @@ -93,68 +58,25 @@ class LuxTTSBackend: await asyncio.to_thread(self._load_model_sync) 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" - is_cached = self._is_model_cached() - # Set up HF progress tracking (intercepts tqdm for file-level progress) - 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: + with model_load_progress(model_name, is_cached): from zipvoice.luxvoice import LuxTTS device = self.device logger.info(f"Loading LuxTTS on {device}...") - # LuxTTS constructor downloads model and loads everything - try: - if device == "cpu": - import os - threads = os.cpu_count() or 4 - self.model = LuxTTS( - model_path=LUXTTS_HF_REPO, - device="cpu", - threads=min(threads, 8), - ) - else: - self.model = LuxTTS( - model_path=LUXTTS_HF_REPO, - device=device, - ) - finally: - tracker_context.__exit__(None, None, None) + if device == "cpu": + import os + threads = os.cpu_count() or 4 + self.model = LuxTTS( + model_path=LUXTTS_HF_REPO, device="cpu", threads=min(threads, 8), + ) + else: + self.model = LuxTTS(model_path=LUXTTS_HF_REPO, device=device) - if not is_cached: - 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 + logger.info("LuxTTS loaded successfully") def unload_model(self) -> None: """Unload model to free memory.""" @@ -205,28 +127,8 @@ class LuxTTSBackend: return encoded, False - async def combine_voice_prompts( - self, - 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 combine_voice_prompts(self, audio_paths, reference_texts): + return await _combine_voice_prompts(audio_paths, reference_texts, sample_rate=24000) async def generate( self, diff --git a/backend/backends/mlx_backend.py b/backend/backends/mlx_backend.py index 617e61b0..1ea7746f 100644 --- a/backend/backends/mlx_backend.py +++ b/backend/backends/mlx_backend.py @@ -14,18 +14,9 @@ from ..utils.hf_offline_patch import patch_huggingface_hub_offline, ensure_origi patch_huggingface_hub_offline() 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.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: @@ -66,45 +57,10 @@ class MLXTTSBackend: return hf_model_id def _is_model_cached(self, model_size: str) -> bool: - """ - 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")) 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 + return is_model_cached( + self._get_model_path(model_size), + weight_extensions=(".safetensors", ".bin", ".npz"), + ) 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): """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: - # Get model path BEFORE importing mlx_audio - model_path = self._get_model_path(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) + with model_load_progress(model_name, is_cached): + from mlx_audio.tts import load + print(f"Loading MLX TTS model {model_size}...") - # Initialize progress state so SSE endpoint has initial data to send - # 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) + try: self.model = load(model_path) - else: - raise - finally: - # Exit the patch context - tracker_context.__exit__(None, None, None) - # Restore original HF_HUB_OFFLINE setting - if original_hf_hub_offline is not None: - os.environ["HF_HUB_OFFLINE"] = original_hf_hub_offline - else: - os.environ.pop("HF_HUB_OFFLINE", 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"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 + except Exception as load_error: + 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) + else: + raise + finally: + if original_hf_hub_offline is not None: + os.environ["HF_HUB_OFFLINE"] = original_hf_hub_offline + else: + os.environ.pop("HF_HUB_OFFLINE", None) + + self._current_model_size = model_size + self.model_size = model_size + print(f"MLX TTS model {model_size} loaded successfully") def unload_model(self): """Unload the model to free memory.""" @@ -288,36 +181,8 @@ class MLXTTSBackend: return voice_prompt_items, False - async def combine_voice_prompts( - self, - 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 combine_voice_prompts(self, audio_paths, reference_texts): + return await _combine_voice_prompts(audio_paths, reference_texts) async def generate( self, @@ -413,14 +278,6 @@ class MLXTTSBackend: 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: """MLX-based STT backend using mlx-audio Whisper.""" @@ -433,45 +290,8 @@ class MLXSTTBackend: return self.model is not None def _is_model_cached(self, model_size: str) -> bool: - """ - Check if the Whisper 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 - 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 + hf_repo = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}") + return is_model_cached(hf_repo, weight_extensions=(".safetensors", ".bin", ".npz")) 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): """Synchronous model loading.""" - try: - progress_manager = get_progress_manager() - task_manager = get_task_manager() - progress_model_name = f"whisper-{model_size}" - - # 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 + progress_model_name = f"whisper-{model_size}" + is_cached = self._is_model_cached(model_size) + + with model_load_progress(progress_model_name, is_cached): 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}") - print(f"Loading MLX Whisper model {model_size}...") - - # Only track download progress if model is NOT cached - if not is_cached: - # Start tracking download task - 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 + self.model = load(model_name) + + self.model_size = model_size + print(f"MLX Whisper model {model_size} loaded successfully") def unload_model(self): """Unload the model to free memory.""" diff --git a/backend/backends/pytorch_backend.py b/backend/backends/pytorch_backend.py index 517403c7..809f6984 100644 --- a/backend/backends/pytorch_backend.py +++ b/backend/backends/pytorch_backend.py @@ -6,20 +6,11 @@ from typing import Optional, List, Tuple import asyncio import torch 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.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", -} +from ..utils.audio import load_audio class PyTorchTTSBackend: @@ -33,26 +24,7 @@ class PyTorchTTSBackend: def _get_device(self) -> str: """Get the best available device.""" - if torch.cuda.is_available(): - 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" + return get_torch_device(allow_xpu=True, allow_directml=True) def is_loaded(self) -> bool: """Check if model is loaded.""" @@ -79,44 +51,7 @@ class PyTorchTTSBackend: return hf_model_map[model_size] def _is_model_cached(self, model_size: str) -> bool: - """ - 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 + return is_model_cached(self._get_model_path(model_size)) 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): """Synchronous model loading.""" - try: - progress_manager = get_progress_manager() - task_manager = get_task_manager() - model_name = f"qwen-tts-{model_size}" + model_name = f"qwen-tts-{model_size}" + is_cached = self._is_model_cached(model_size) - # 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 (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 + with model_load_progress(model_name, is_cached): from qwen_tts import Qwen3TTSModel - - # Get model path (local or HuggingFace Hub ID) model_path = self._get_model_path(model_size) - print(f"Loading TTS model {model_size} on {self.device}...") - # 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 - 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", + 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, ) - # Load the model (tqdm is patched, but filters out non-download progress) - try: - # Don't pass device_map on CPU: accelerate's meta-tensor mechanism - # 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 + self._current_model_size = model_size + self.model_size = model_size + print(f"TTS model {model_size} loaded successfully") def unload_model(self): """Unload the model to free memory.""" @@ -303,31 +174,7 @@ class PyTorchTTSBackend: 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 + return await _combine_voice_prompts(audio_paths, reference_texts) async def generate( self, @@ -376,15 +223,6 @@ class PyTorchTTSBackend: 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: """PyTorch-based STT backend using Whisper.""" @@ -396,69 +234,15 @@ class PyTorchSTTBackend: def _get_device(self) -> str: """Get the best available device.""" - if torch.cuda.is_available(): - 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" + return get_torch_device(allow_xpu=True, allow_directml=True) def is_loaded(self) -> bool: """Check if model is loaded.""" return self.model is not None def _is_model_cached(self, model_size: str) -> bool: - """ - Check if the Whisper 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 - 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 + hf_repo = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}") + return is_model_cached(hf_repo) async def load_model_async(self, model_size: Optional[str] = None): """ @@ -467,94 +251,33 @@ class PyTorchSTTBackend: Args: 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: 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: - print(f"[DEBUG] Early return - model already loaded") 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) - print(f"[DEBUG] asyncio.to_thread completed") # Alias for compatibility load_model = load_model_async def _load_model_sync(self, model_size: str): """Synchronous model loading.""" - print(f"[DEBUG] _load_model_sync called for Whisper {model_size}") - try: - progress_manager = get_progress_manager() - task_manager = get_task_manager() - progress_model_name = f"whisper-{model_size}" + progress_model_name = f"whisper-{model_size}" + is_cached = self._is_model_cached(model_size) - # 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 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 + with model_load_progress(progress_model_name, is_cached): from transformers import WhisperProcessor, WhisperForConditionalGeneration - 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}...") - # Only track download progress if model is NOT cached - if not is_cached: - # Start tracking download task - task_manager.start_download(progress_model_name) + self.processor = WhisperProcessor.from_pretrained(model_name) + self.model = WhisperForConditionalGeneration.from_pretrained(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, # 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 + self.model.to(self.device) + self.model_size = model_size + print(f"Whisper model {model_size} loaded successfully") def unload_model(self): """Unload the model to free memory.""" diff --git a/backend/export_import.py b/backend/export_import.py index 7af1f9b0..58acd400 100644 --- a/backend/export_import.py +++ b/backend/export_import.py @@ -19,11 +19,6 @@ from .models import VoiceProfileCreate 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: """ 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 samples_data = {} - profile_dir = _get_profiles_dir() / profile_id + profile_dir = config.get_profiles_dir() / profile_id for sample in samples: # 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) # 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) # Handle avatar if present diff --git a/backend/history.py b/backend/history.py index 78da51ca..c2c9197a 100644 --- a/backend/history.py +++ b/backend/history.py @@ -15,11 +15,6 @@ from .database import Generation as DBGeneration, GenerationVersion as DBGenerat 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: """Get versions list and active version ID for a generation.""" import json diff --git a/backend/main.py b/backend/main.py index 5c9fca44..c16a85b3 100644 --- a/backend/main.py +++ b/backend/main.py @@ -18,7 +18,6 @@ import tempfile import io from pathlib import Path import uuid -import asyncio import signal 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 .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.tasks import get_task_manager from .utils.cache import clear_voice_prompt_cache @@ -1980,31 +1980,6 @@ async def update_profile_effects( 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 # ============================================ diff --git a/backend/migrate_add_instruct.py b/backend/migrate_add_instruct.py deleted file mode 100644 index 4d899cf3..00000000 --- a/backend/migrate_add_instruct.py +++ /dev/null @@ -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() diff --git a/backend/profiles.py b/backend/profiles.py index c7a48661..9f8f3b5d 100644 --- a/backend/profiles.py +++ b/backend/profiles.py @@ -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( data: VoiceProfileCreate, db: Session, @@ -100,7 +95,7 @@ async def create_profile( db.refresh(db_profile) # 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) return _profile_to_response(db_profile) @@ -136,7 +131,7 @@ async def add_profile_sample( # Create sample ID and directory 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) # Copy audio file to profile directory @@ -316,7 +311,7 @@ async def delete_profile( db.commit() # Delete profile directory - profile_dir = _get_profiles_dir() / profile_id + profile_dir = config.get_profiles_dir() / profile_id if profile_dir.exists(): shutil.rmtree(profile_dir) @@ -516,7 +511,7 @@ async def upload_avatar( ext = ext_map.get(img_format, '.png') # 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) output_path = profile_dir / f"avatar{ext}" diff --git a/backend/studio.py b/backend/studio.py deleted file mode 100644 index 027a9b65..00000000 --- a/backend/studio.py +++ /dev/null @@ -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") diff --git a/backend/tests/test_profile_duplicate_names.py b/backend/tests/test_profile_duplicate_names.py index 81d4f7ef..55ee8587 100644 --- a/backend/tests/test_profile_duplicate_names.py +++ b/backend/tests/test_profile_duplicate_names.py @@ -45,8 +45,8 @@ def test_db(): @pytest.fixture def mock_profiles_dir(monkeypatch, tmp_path): """Mock the profiles directory to use a temporary path.""" - import profiles - monkeypatch.setattr(profiles, '_get_profiles_dir', lambda: tmp_path) + from backend import config + monkeypatch.setattr(config, 'get_profiles_dir', lambda: tmp_path) return tmp_path diff --git a/backend/utils/validation.py b/backend/utils/validation.py deleted file mode 100644 index 3637da86..00000000 --- a/backend/utils/validation.py +++ /dev/null @@ -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