mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 01:25:18 -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,
|
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)
|
||||||
|
|||||||
Reference in New Issue
Block a user