diff --git a/backend/backends/chatterbox_mlx_backend.py b/backend/backends/chatterbox_mlx_backend.py index 57a88224..d4629426 100644 --- a/backend/backends/chatterbox_mlx_backend.py +++ b/backend/backends/chatterbox_mlx_backend.py @@ -25,6 +25,7 @@ from .base import ( is_model_cached, model_load_progress, ) +from .mlx_backend import _run_on_mlx_thread logger = logging.getLogger(__name__) @@ -40,7 +41,8 @@ class ChatterboxMLXTTSBackend: def __init__(self): self.model = None self.model_size = "default" - self._model_load_lock = asyncio.Lock() + # Guards the load-then-use sequence, as MLXTTSBackend._op_lock does. + self._op_lock = asyncio.Lock() def is_loaded(self) -> bool: return self.model is not None @@ -55,10 +57,13 @@ class ChatterboxMLXTTSBackend: """Load the Chatterbox multilingual MLX model.""" if self.model is not None: return - async with self._model_load_lock: + async with self._op_lock: if self.model is not None: return - await asyncio.to_thread(self._load_model_sync) + # MLX streams are thread-local: every MLX call in the process + # shares the single worker in mlx_backend (issue #699), so load + # and generate always land on the same OS thread. + await _run_on_mlx_thread(self._load_model_sync) def _load_model_sync(self): """Synchronous model loading.""" @@ -145,8 +150,6 @@ class ChatterboxMLXTTSBackend: 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}") @@ -155,6 +158,13 @@ class ChatterboxMLXTTSBackend: def _generate_sync(): import mlx.core as mx # lazy: heavy import + # Load (if needed) and generate as ONE worker submission, as in + # MLXTTSBackend._reload_and_generate_sync, so an unload cannot + # land between the two; bind the model locally for the same reason. + if self.model is None: + self._load_model_sync() + model = self.model + if seed is not None: mx.random.seed(seed) @@ -163,7 +173,7 @@ class ChatterboxMLXTTSBackend: # mlx-audio yields GenerationResult chunks; the whole clip is their concatenation. chunks = [ np.asarray(result.audio).squeeze() - for result in self.model.generate( + for result in model.generate( text, ref_audio=ref_audio, lang_code=language, @@ -173,7 +183,8 @@ class ChatterboxMLXTTSBackend: ] 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 + sample_rate = getattr(model, "sr", None) or 24000 return audio, int(sample_rate) - return await asyncio.to_thread(_generate_sync) + async with self._op_lock: + return await _run_on_mlx_thread(_generate_sync)