From ad5d64c0a3d25e6172c69bb768fd169f6bd182c2 Mon Sep 17 00:00:00 2001 From: Charles Hasse Date: Tue, 11 Aug 2026 16:40:56 -0300 Subject: [PATCH] feat(backend): run Chatterbox multilingual on MLX for Apple Silicon The Chatterbox backend is pinned to the CPU on macOS, so voice cloning runs at roughly 4x realtime there. This adds an MLX/Metal backend for the same engine and selects it on Apple Silicon, mirroring the split the qwen engine already makes between mlx_backend and pytorch_backend. Measured on an M4 Max (36 GB) with a cloned pt-BR profile, same API, same profile, model already loaded: short sentence (1.8s of audio): 7.4-9.2s -> 1.2s longer sentence (5.0s of audio): 23.7s -> 3.1s The model config for chatterbox-tts is now backend aware, same as the qwen configs, so the download matches the backend that will consume it. Nothing changes off Apple Silicon: the PyTorch backend is still selected there, and the CPU pinning it relies on is untouched. --- backend/backends/__init__.py | 28 +++- backend/backends/chatterbox_mlx_backend.py | 179 +++++++++++++++++++++ 2 files changed, 200 insertions(+), 7 deletions(-) create mode 100644 backend/backends/chatterbox_mlx_backend.py diff --git a/backend/backends/__init__.py b/backend/backends/__init__.py index e3cfe049..9c539da0 100644 --- a/backend/backends/__init__.py +++ b/backend/backends/__init__.py @@ -290,10 +290,17 @@ def _get_qwen_custom_voice_configs() -> list[ModelConfig]: def _get_non_qwen_tts_configs() -> list[ModelConfig]: - """Return model configs for non-Qwen TTS engines. + """Return model configs for non-Qwen TTS engines.""" + # Chatterbox multilingual follows the same backend-aware split as Qwen: the MLX + # backend loads pre-converted weights, so the download must match the backend that + # will consume it. + if get_backend_type() == "mlx": + chatterbox_repo = "mlx-community/chatterbox-multilingual-v3" + chatterbox_size_mb = 2600 + else: + chatterbox_repo = "ResembleAI/chatterbox" + chatterbox_size_mb = 3200 - These are static — no backend-type branching needed. - """ return [ ModelConfig( model_name="luxtts", @@ -307,8 +314,8 @@ def _get_non_qwen_tts_configs() -> list[ModelConfig]: model_name="chatterbox-tts", display_name="Chatterbox TTS (Multilingual)", engine="chatterbox", - hf_repo_id="ResembleAI/chatterbox", - size_mb=3200, + hf_repo_id=chatterbox_repo, + size_mb=chatterbox_size_mb, needs_trim=True, languages=[ "zh", @@ -707,9 +714,16 @@ def get_tts_backend_for_engine(engine: str) -> TTSBackend: backend = LuxTTSBackend() elif engine == "chatterbox": - from .chatterbox_backend import ChatterboxTTSBackend + # Same split the qwen engine already makes: on Apple Silicon the MLX/Metal + # port renders 7-9x faster than the CPU-pinned PyTorch path. + if get_backend_type() == "mlx": + from .chatterbox_mlx_backend import ChatterboxMLXTTSBackend - backend = ChatterboxTTSBackend() + backend = ChatterboxMLXTTSBackend() + else: + from .chatterbox_backend import ChatterboxTTSBackend + + backend = ChatterboxTTSBackend() elif engine == "chatterbox_turbo": from .chatterbox_turbo_backend import ChatterboxTurboTTSBackend diff --git a/backend/backends/chatterbox_mlx_backend.py b/backend/backends/chatterbox_mlx_backend.py new file mode 100644 index 00000000..57a88224 --- /dev/null +++ b/backend/backends/chatterbox_mlx_backend.py @@ -0,0 +1,179 @@ +""" +Chatterbox multilingual TTS backend, MLX (Metal) flavour. + +Same zero-shot voice cloning and same 23 languages as ``chatterbox_backend``, but running +on the Apple Silicon GPU through mlx-audio instead of PyTorch on the CPU. + +The PyTorch path is pinned to the CPU on macOS (see ``chatterbox_backend``), which costs +roughly 4x realtime. This backend uses the pre-converted weights published as +``mlx-community/chatterbox-multilingual-v3`` and renders the same sentences at about 0.5x +realtime on an M4 Max with a cloned pt-BR profile. + +This mirrors the split the qwen engine already makes between ``mlx_backend`` and +``pytorch_backend``: MLX where it is available, PyTorch everywhere else. +""" + +import asyncio +import logging +from pathlib import Path +from typing import ClassVar + +import numpy as np + +from .base import ( + combine_voice_prompts as _combine_voice_prompts, + is_model_cached, + model_load_progress, +) + +logger = logging.getLogger(__name__) + +CHATTERBOX_MLX_HF_REPO = "mlx-community/chatterbox-multilingual-v3" + +# Files that must be present for the MLX multilingual model +_MLX_WEIGHT_FILES = ["model.safetensors", "config.json", "tokenizer.json"] + + +class ChatterboxMLXTTSBackend: + """Chatterbox Multilingual TTS backend for voice cloning, on MLX/Metal.""" + + def __init__(self): + self.model = None + self.model_size = "default" + self._model_load_lock = asyncio.Lock() + + def is_loaded(self) -> bool: + return self.model is not None + + def _get_model_path(self, model_size: str = "default") -> str: + return CHATTERBOX_MLX_HF_REPO + + def _is_model_cached(self, model_size: str = "default") -> bool: + return is_model_cached(CHATTERBOX_MLX_HF_REPO, required_files=_MLX_WEIGHT_FILES) + + async def load_model(self, model_size: str = "default") -> None: + """Load the Chatterbox multilingual MLX model.""" + if self.model is not None: + return + async with self._model_load_lock: + if self.model is not None: + return + await asyncio.to_thread(self._load_model_sync) + + def _load_model_sync(self): + """Synchronous model loading.""" + is_cached = self._is_model_cached() + + with model_load_progress("chatterbox-tts", is_cached): + from huggingface_hub import snapshot_download # lazy: heavy import + from mlx_audio.tts.models.chatterbox.chatterbox import Model # lazy: heavy import + + logger.info("Loading Chatterbox Multilingual TTS on MLX (Metal)...") + ckpt_dir = snapshot_download(CHATTERBOX_MLX_HF_REPO) + self.model = Model.from_pretrained(ckpt_dir) + + logger.info("Chatterbox Multilingual TTS (MLX) loaded successfully") + + def unload_model(self) -> None: + """Unload model to free memory.""" + if self.model is None: + return + + del self.model + self.model = None + try: + import mlx.core as mx # lazy: heavy import + + mx.clear_cache() + except Exception: + logger.debug("mlx cache not cleared", exc_info=True) + logger.info("Chatterbox (MLX) unloaded") + + async def create_voice_prompt( + self, + audio_path: str, + reference_text: str, + use_cache: bool = True, + ) -> tuple[dict, bool]: + """ + Create voice prompt from reference audio. + + Chatterbox conditions on the reference audio at generation time, so the prompt + just stores the file path. + """ + voice_prompt = { + "ref_audio": str(audio_path), + "ref_text": reference_text, + } + return voice_prompt, False + + async def combine_voice_prompts( + self, + audio_paths: list[str], + reference_texts: list[str], + ) -> tuple[np.ndarray, str]: + return await _combine_voice_prompts(audio_paths, reference_texts) + + # The MLX port carries its own sampling defaults, validated by ear against the PyTorch + # output on a cloned profile. The PyTorch tuning (repetition_penalty=2.0) is not + # transferable: the two implementations weight it differently. + _DEFAULTS: ClassVar[dict] = { + "exaggeration": 0.1, + "cfg_weight": 0.5, + "temperature": 0.8, + "repetition_penalty": 1.2, + } + + async def generate( + self, + text: str, + voice_prompt: dict, + language: str = "en", + seed: int | None = None, + instruct: str | None = None, + ) -> tuple[np.ndarray, int]: + """ + Generate audio using Chatterbox Multilingual TTS on MLX. + + Args: + text: Text to synthesize + voice_prompt: Dict with ref_audio path + language: BCP-47 language code + seed: Random seed for reproducibility + instruct: Unused (protocol compatibility) + + Returns: + Tuple of (audio_array, sample_rate) + """ + await self.load_model() + + ref_audio = voice_prompt.get("ref_audio") + if ref_audio and not Path(ref_audio).exists(): + logger.warning(f"Reference audio not found: {ref_audio}") + ref_audio = None + + def _generate_sync(): + import mlx.core as mx # lazy: heavy import + + if seed is not None: + mx.random.seed(seed) + + logger.info(f"[Chatterbox/MLX] Generating: lang={language}") + + # mlx-audio yields GenerationResult chunks; the whole clip is their concatenation. + chunks = [ + np.asarray(result.audio).squeeze() + for result in self.model.generate( + text, + ref_audio=ref_audio, + lang_code=language, + verbose=False, + **self._DEFAULTS, + ) + ] + audio = np.concatenate(chunks).astype(np.float32) if chunks else np.zeros(0, dtype=np.float32) + + sample_rate = getattr(self.model, "sr", None) or 24000 + return audio, int(sample_rate) + + return await asyncio.to_thread(_generate_sync)