fix(backend): keep the trimmed clip when a runaway retry cannot split further; track the S3 tokenizer repo in model status and delete

Review follow-ups: with retries_runaway on for an engine that also has a
trim step, a <=100-char chunk flagged as runaway raised instead of
falling back to the trimmed output that already cuts the silence+noise
tail; generate_chunked now returns the trimmed chunk in that terminal
case. ModelConfig gains aux_hf_repo_ids so /models status only reports
the Chatterbox MLX model as downloaded once mlx-community/S3TokenizerV2
is present too (matching the backend's own cache check) and
DELETE /models/{name} removes that repo as well.
This commit is contained in:
jamiepine
2026-10-04 00:25:58 +00:00
committed by capy-ai-staging[bot]
parent ad6ec3c6ef
commit a2b453afed
4 changed files with 40 additions and 2 deletions
+11
View File
@@ -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 @dataclass
class ModelConfig: class ModelConfig:
"""Declarative config for a downloadable model variant.""" """Declarative config for a downloadable model variant."""
@@ -53,6 +58,9 @@ class ModelConfig:
display_name: str # e.g. "LuxTTS (Fast, CPU-friendly)" display_name: str # e.g. "LuxTTS (Fast, CPU-friendly)"
engine: str # e.g. "luxtts", "chatterbox" engine: str # e.g. "luxtts", "chatterbox"
hf_repo_id: str # e.g. "YatharthS/LuxTTS" 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" model_size: str = "default"
size_mb: int = 0 size_mb: int = 0
needs_trim: bool = False needs_trim: bool = False
@@ -299,9 +307,11 @@ def _get_non_qwen_tts_configs() -> list[ModelConfig]:
chatterbox_repo = "mlx-community/chatterbox-multilingual-v3" chatterbox_repo = "mlx-community/chatterbox-multilingual-v3"
# 2.5 GB of weights plus the separately fetched S3TokenizerV2 (~470 MB) # 2.5 GB of weights plus the separately fetched S3TokenizerV2 (~470 MB)
chatterbox_size_mb = 3000 chatterbox_size_mb = 3000
chatterbox_aux_repos = (CHATTERBOX_MLX_S3_TOKENIZER_REPO,)
else: else:
chatterbox_repo = "ResembleAI/chatterbox" chatterbox_repo = "ResembleAI/chatterbox"
chatterbox_size_mb = 3200 chatterbox_size_mb = 3200
chatterbox_aux_repos = ()
return [ return [
ModelConfig( ModelConfig(
@@ -317,6 +327,7 @@ def _get_non_qwen_tts_configs() -> list[ModelConfig]:
display_name="Chatterbox TTS (Multilingual)", display_name="Chatterbox TTS (Multilingual)",
engine="chatterbox", engine="chatterbox",
hf_repo_id=chatterbox_repo, hf_repo_id=chatterbox_repo,
aux_hf_repo_ids=chatterbox_aux_repos,
size_mb=chatterbox_size_mb, size_mb=chatterbox_size_mb,
needs_trim=True, needs_trim=True,
# Same EOS miss the qwen configs guard against: on mlx-audio the decoder can run past # Same EOS miss the qwen configs guard against: on mlx-audio the decoder can run past
+2 -1
View File
@@ -20,6 +20,7 @@ from typing import ClassVar
import numpy as np import numpy as np
from . import CHATTERBOX_MLX_S3_TOKENIZER_REPO
from .base import ( from .base import (
combine_voice_prompts as _combine_voice_prompts, combine_voice_prompts as _combine_voice_prompts,
is_model_cached, is_model_cached,
@@ -32,7 +33,7 @@ logger = logging.getLogger(__name__)
CHATTERBOX_MLX_HF_REPO = "mlx-community/chatterbox-multilingual-v3" CHATTERBOX_MLX_HF_REPO = "mlx-community/chatterbox-multilingual-v3"
# mlx-audio's Model.from_pretrained fetches the S3 speech tokenizer from this # 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. # 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 # Files that must be present for the MLX multilingual model
_MLX_WEIGHT_FILES = ["model.safetensors", "config.json", "tokenizer.json"] _MLX_WEIGHT_FILES = ["model.safetensors", "config.json", "tokenizer.json"]
+18 -1
View File
@@ -306,7 +306,7 @@ async def get_model_status():
use_scan_cache = False use_scan_cache = False
from ..backends import get_all_model_configs, check_model_loaded 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() registry_configs = get_all_model_configs()
model_configs = [ model_configs = [
@@ -314,6 +314,7 @@ async def get_model_status():
"model_name": cfg.model_name, "model_name": cfg.model_name,
"display_name": cfg.display_name, "display_name": cfg.display_name,
"hf_repo_id": cfg.hf_repo_id, "hf_repo_id": cfg.hf_repo_id,
"aux_hf_repo_ids": cfg.aux_hf_repo_ids,
"model_size": cfg.model_size, "model_size": cfg.model_size,
"check_loaded": lambda c=cfg: check_model_loaded(c), "check_loaded": lambda c=cfg: check_model_loaded(c),
} }
@@ -404,6 +405,12 @@ async def get_model_status():
except Exception: except Exception:
pass 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: try:
loaded = config["check_loaded"]() loaded = config["check_loaded"]()
except Exception: except Exception:
@@ -427,6 +434,12 @@ async def get_model_status():
) )
) )
except Exception: 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: try:
loaded = config["check_loaded"]() loaded = config["check_loaded"]()
except Exception: except Exception:
@@ -529,6 +542,10 @@ async def delete_model(model_name: str):
try: try:
shutil.rmtree(repo_cache_dir) 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: except OSError as e:
raise HTTPException(status_code=500, detail=f"Failed to delete model cache directory: {str(e)}") raise HTTPException(status_code=500, detail=f"Failed to delete model cache directory: {str(e)}")
+9
View File
@@ -265,6 +265,15 @@ async def generate_chunked(
if runaway_detector is not None and runaway_detector(chunk_audio, chunk_sr): 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 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( raise RuntimeError(
"TTS output remained unstable after retrying smaller text chunks" "TTS output remained unstable after retrying smaller text chunks"
) )