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:
Jamie Pine
2026-03-16 01:10:02 -07:00
parent 9514c6596c
commit 0813a3d9d6
16 changed files with 480 additions and 1281 deletions
+31 -20
View File
@@ -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
+14
View File
@@ -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:
+277
View File
@@ -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)
+28 -159
View File
@@ -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] = {
+20 -156
View File
@@ -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,
+19 -117
View File
@@ -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
View File
@@ -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."""
+34 -311
View File
@@ -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."""
+2 -7
View File
@@ -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
-5
View File
@@ -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
View File
@@ -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
# ============================================ # ============================================
-48
View File
@@ -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
View File
@@ -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}"
-66
View File
@@ -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
-66
View File
@@ -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