fix(backend): run Chatterbox MLX load and generate on the shared MLX worker thread

With asyncio.to_thread the new backend hits the same thread-local stream
crash as the Qwen MLX backend once the default pool is busy (reproduced
on an M2 Ultra: 'There is no Stream(cpu, 2) in current thread.' on the
first generate after load). Route load and generate through
mlx_backend._run_on_mlx_thread, fold load-if-needed + generate into one
worker submission under the backend's lock, and bind the model locally
inside the closure, matching MLXTTSBackend.
This commit is contained in:
jamiepine
2026-10-04 00:25:58 +00:00
committed by capy-ai-staging[bot]
parent 502123a943
commit c5cf7436f7
+19 -8
View File
@@ -25,6 +25,7 @@ from .base import (
is_model_cached, is_model_cached,
model_load_progress, model_load_progress,
) )
from .mlx_backend import _run_on_mlx_thread
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -40,7 +41,8 @@ class ChatterboxMLXTTSBackend:
def __init__(self): def __init__(self):
self.model = None self.model = None
self.model_size = "default" 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: def is_loaded(self) -> bool:
return self.model is not None return self.model is not None
@@ -55,10 +57,13 @@ class ChatterboxMLXTTSBackend:
"""Load the Chatterbox multilingual MLX model.""" """Load the Chatterbox multilingual MLX model."""
if self.model is not None: if self.model is not None:
return return
async with self._model_load_lock: async with self._op_lock:
if self.model is not None: if self.model is not None:
return 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): def _load_model_sync(self):
"""Synchronous model loading.""" """Synchronous model loading."""
@@ -145,8 +150,6 @@ class ChatterboxMLXTTSBackend:
Returns: Returns:
Tuple of (audio_array, sample_rate) Tuple of (audio_array, sample_rate)
""" """
await self.load_model()
ref_audio = voice_prompt.get("ref_audio") ref_audio = voice_prompt.get("ref_audio")
if ref_audio and not Path(ref_audio).exists(): if ref_audio and not Path(ref_audio).exists():
logger.warning(f"Reference audio not found: {ref_audio}") logger.warning(f"Reference audio not found: {ref_audio}")
@@ -155,6 +158,13 @@ class ChatterboxMLXTTSBackend:
def _generate_sync(): def _generate_sync():
import mlx.core as mx # lazy: heavy import 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: if seed is not None:
mx.random.seed(seed) mx.random.seed(seed)
@@ -163,7 +173,7 @@ class ChatterboxMLXTTSBackend:
# mlx-audio yields GenerationResult chunks; the whole clip is their concatenation. # mlx-audio yields GenerationResult chunks; the whole clip is their concatenation.
chunks = [ chunks = [
np.asarray(result.audio).squeeze() np.asarray(result.audio).squeeze()
for result in self.model.generate( for result in model.generate(
text, text,
ref_audio=ref_audio, ref_audio=ref_audio,
lang_code=language, 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) 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 audio, int(sample_rate)
return await asyncio.to_thread(_generate_sync) async with self._op_lock:
return await _run_on_mlx_thread(_generate_sync)