mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 09:35:16 -07:00
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:
committed by
capy-ai-staging[bot]
parent
502123a943
commit
c5cf7436f7
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user