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.
This commit is contained in:
Charles Hasse
2026-10-04 00:25:58 +00:00
committed by capy-ai-staging[bot]
parent 63899fd865
commit ad5d64c0a3
2 changed files with 200 additions and 7 deletions
+19 -5
View File
@@ -290,10 +290,17 @@ def _get_qwen_custom_voice_configs() -> list[ModelConfig]:
def _get_non_qwen_tts_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 [ return [
ModelConfig( ModelConfig(
model_name="luxtts", model_name="luxtts",
@@ -307,8 +314,8 @@ def _get_non_qwen_tts_configs() -> list[ModelConfig]:
model_name="chatterbox-tts", model_name="chatterbox-tts",
display_name="Chatterbox TTS (Multilingual)", display_name="Chatterbox TTS (Multilingual)",
engine="chatterbox", engine="chatterbox",
hf_repo_id="ResembleAI/chatterbox", hf_repo_id=chatterbox_repo,
size_mb=3200, size_mb=chatterbox_size_mb,
needs_trim=True, needs_trim=True,
languages=[ languages=[
"zh", "zh",
@@ -707,6 +714,13 @@ def get_tts_backend_for_engine(engine: str) -> TTSBackend:
backend = LuxTTSBackend() backend = LuxTTSBackend()
elif engine == "chatterbox": elif engine == "chatterbox":
# 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 = ChatterboxMLXTTSBackend()
else:
from .chatterbox_backend import ChatterboxTTSBackend from .chatterbox_backend import ChatterboxTTSBackend
backend = ChatterboxTTSBackend() backend = ChatterboxTTSBackend()
+179
View File
@@ -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)