diff --git a/backend/backends/__init__.py b/backend/backends/__init__.py index 5d006d59..4da1a439 100644 --- a/backend/backends/__init__.py +++ b/backend/backends/__init__.py @@ -45,6 +45,11 @@ WHISPER_HF_REPOS = { } +# mlx-audio's Chatterbox loader fetches the S3 speech tokenizer from this +# second repo; see chatterbox_mlx_backend. +CHATTERBOX_MLX_S3_TOKENIZER_REPO = "mlx-community/S3TokenizerV2" + + @dataclass class ModelConfig: """Declarative config for a downloadable model variant.""" @@ -53,6 +58,9 @@ class ModelConfig: display_name: str # e.g. "LuxTTS (Fast, CPU-friendly)" engine: str # e.g. "luxtts", "chatterbox" hf_repo_id: str # e.g. "YatharthS/LuxTTS" + # Extra HF repos the backend fetches at load time (e.g. a shared + # tokenizer); download status and delete must account for them too. + aux_hf_repo_ids: tuple[str, ...] = () model_size: str = "default" size_mb: int = 0 needs_trim: bool = False @@ -299,9 +307,11 @@ def _get_non_qwen_tts_configs() -> list[ModelConfig]: chatterbox_repo = "mlx-community/chatterbox-multilingual-v3" # 2.5 GB of weights plus the separately fetched S3TokenizerV2 (~470 MB) chatterbox_size_mb = 3000 + chatterbox_aux_repos = (CHATTERBOX_MLX_S3_TOKENIZER_REPO,) else: chatterbox_repo = "ResembleAI/chatterbox" chatterbox_size_mb = 3200 + chatterbox_aux_repos = () return [ ModelConfig( @@ -317,6 +327,7 @@ def _get_non_qwen_tts_configs() -> list[ModelConfig]: display_name="Chatterbox TTS (Multilingual)", engine="chatterbox", hf_repo_id=chatterbox_repo, + aux_hf_repo_ids=chatterbox_aux_repos, size_mb=chatterbox_size_mb, needs_trim=True, # Same EOS miss the qwen configs guard against: on mlx-audio the decoder can run past diff --git a/backend/backends/chatterbox_mlx_backend.py b/backend/backends/chatterbox_mlx_backend.py index ca26c75a..5d59f2f8 100644 --- a/backend/backends/chatterbox_mlx_backend.py +++ b/backend/backends/chatterbox_mlx_backend.py @@ -20,6 +20,7 @@ from typing import ClassVar import numpy as np +from . import CHATTERBOX_MLX_S3_TOKENIZER_REPO from .base import ( combine_voice_prompts as _combine_voice_prompts, is_model_cached, @@ -32,7 +33,7 @@ logger = logging.getLogger(__name__) CHATTERBOX_MLX_HF_REPO = "mlx-community/chatterbox-multilingual-v3" # mlx-audio's Model.from_pretrained fetches the S3 speech tokenizer from this # second repo (~470 MB), so the engine is only "downloaded" once both are cached. -S3_TOKENIZER_HF_REPO = "mlx-community/S3TokenizerV2" +S3_TOKENIZER_HF_REPO = CHATTERBOX_MLX_S3_TOKENIZER_REPO # Files that must be present for the MLX multilingual model _MLX_WEIGHT_FILES = ["model.safetensors", "config.json", "tokenizer.json"] diff --git a/backend/routes/models.py b/backend/routes/models.py index b7533465..e00e71ff 100644 --- a/backend/routes/models.py +++ b/backend/routes/models.py @@ -306,7 +306,7 @@ async def get_model_status(): use_scan_cache = False from ..backends import get_all_model_configs, check_model_loaded - from ..backends.base import has_in_progress_download + from ..backends.base import has_in_progress_download, is_model_cached registry_configs = get_all_model_configs() model_configs = [ @@ -314,6 +314,7 @@ async def get_model_status(): "model_name": cfg.model_name, "display_name": cfg.display_name, "hf_repo_id": cfg.hf_repo_id, + "aux_hf_repo_ids": cfg.aux_hf_repo_ids, "model_size": cfg.model_size, "check_loaded": lambda c=cfg: check_model_loaded(c), } @@ -404,6 +405,12 @@ async def get_model_status(): except Exception: pass + # A model whose backend also pulls auxiliary repos at load time + # (e.g. Chatterbox MLX's S3 tokenizer) is only downloaded once + # those are present too, matching the backend's own cache check. + if downloaded and not all(is_model_cached(repo) for repo in config["aux_hf_repo_ids"]): + downloaded = False + try: loaded = config["check_loaded"]() except Exception: @@ -427,6 +434,12 @@ async def get_model_status(): ) ) except Exception: + # A model whose backend also pulls auxiliary repos at load time + # (e.g. Chatterbox MLX's S3 tokenizer) is only downloaded once + # those are present too, matching the backend's own cache check. + if downloaded and not all(is_model_cached(repo) for repo in config["aux_hf_repo_ids"]): + downloaded = False + try: loaded = config["check_loaded"]() except Exception: @@ -529,6 +542,10 @@ async def delete_model(model_name: str): try: shutil.rmtree(repo_cache_dir) + for aux_repo_id in config.aux_hf_repo_ids: + aux_cache_dir = Path(cache_dir) / ("models--" + aux_repo_id.replace("/", "--")) + if aux_cache_dir.exists(): + shutil.rmtree(aux_cache_dir) except OSError as e: raise HTTPException(status_code=500, detail=f"Failed to delete model cache directory: {str(e)}") diff --git a/backend/utils/chunked_tts.py b/backend/utils/chunked_tts.py index 7dbefdb7..e3ed5ffe 100644 --- a/backend/utils/chunked_tts.py +++ b/backend/utils/chunked_tts.py @@ -265,6 +265,15 @@ async def generate_chunked( if runaway_detector is not None and runaway_detector(chunk_audio, chunk_sr): if retry_depth >= MAX_RUNAWAY_RETRIES or len(chunk_text) <= MIN_RUNAWAY_RETRY_CHARS: + if trim_fn is not None: + # Engines with a trim step (Chatterbox) already cut the + # silence-then-noise tail; prefer the trimmed clip over + # failing the whole generation when we cannot split further. + logger.warning( + "Unstable TTS output for %d chars could not be retried further; keeping trimmed output", + len(chunk_text), + ) + return np.asarray(trim_fn(chunk_audio, chunk_sr), dtype=np.float32), chunk_sr raise RuntimeError( "TTS output remained unstable after retrying smaller text chunks" )