mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 01:25:18 -07:00
fix(mlx): pin all MLX ops to a single worker thread
Qwen3-TTS (and MLX STT) generation crashed with: "There is no Stream(gpu, N) in current thread." MLXTTSBackend/MLXSTTBackend dispatched model load and generate/ transcribe as separate asyncio.to_thread() calls, which round-robin across Python's default multi-worker executor pool. MLX's Metal backend keeps GPU streams registered per-OS-thread, so a model loaded on one worker thread and then used for generation on a different worker thread hits a missing stream and crashes. Reproduced 100% of the time on macOS/Apple Silicon cloning with both the 1.7B and 0.6B Qwen3-TTS models; Chatterbox/Kokoro were unaffected since they use the PyTorch backend, not this module. Fix: route all four MLX call sites in this file (TTS load, TTS generate, STT load, STT transcribe) through a dedicated single-worker ThreadPoolExecutor instead of asyncio.to_thread's shared pool, so every MLX operation for a given process runs on the same OS thread. Verified: direct /generate API calls against both model sizes completed cleanly after the fix (previously failed every time).
This commit is contained in:
committed by
capy-ai-staging[bot]
parent
76dd300f84
commit
f135a471ae
@@ -7,9 +7,24 @@ import asyncio
|
|||||||
import logging
|
import logging
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# MLX's Metal backend keeps a per-thread stream registry. Loading a model on
|
||||||
|
# one worker thread (via asyncio.to_thread, which round-robins across the
|
||||||
|
# default executor's pool) and then generating on a different worker thread
|
||||||
|
# raises "There is no Stream(gpu, N) in current thread." All MLX calls in
|
||||||
|
# this module must therefore run on the SAME OS thread for the process
|
||||||
|
# lifetime — route them through this single-worker executor instead of
|
||||||
|
# asyncio.to_thread's shared multi-worker pool.
|
||||||
|
_mlx_executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="mlx-worker")
|
||||||
|
|
||||||
|
|
||||||
|
def _run_on_mlx_thread(func, *args):
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
return loop.run_in_executor(_mlx_executor, func, *args)
|
||||||
|
|
||||||
# PATCH: Import and apply offline patch BEFORE any huggingface_hub usage
|
# PATCH: Import and apply offline patch BEFORE any huggingface_hub usage
|
||||||
# This prevents mlx_audio from making network requests when models are cached
|
# This prevents mlx_audio from making network requests when models are cached
|
||||||
from ..utils.hf_offline_patch import patch_huggingface_hub_offline, ensure_original_qwen_config_cached
|
from ..utils.hf_offline_patch import patch_huggingface_hub_offline, ensure_original_qwen_config_cached
|
||||||
@@ -82,7 +97,7 @@ class MLXTTSBackend:
|
|||||||
self.unload_model()
|
self.unload_model()
|
||||||
|
|
||||||
# Run blocking load in thread pool
|
# Run blocking load in thread pool
|
||||||
await asyncio.to_thread(self._load_model_sync, model_size)
|
await _run_on_mlx_thread(self._load_model_sync, model_size)
|
||||||
|
|
||||||
# Alias for compatibility
|
# Alias for compatibility
|
||||||
load_model = load_model_async
|
load_model = load_model_async
|
||||||
@@ -259,7 +274,7 @@ class MLXTTSBackend:
|
|||||||
return audio, sample_rate
|
return audio, sample_rate
|
||||||
|
|
||||||
# Run blocking inference in thread pool
|
# Run blocking inference in thread pool
|
||||||
audio, sample_rate = await asyncio.to_thread(_generate_sync)
|
audio, sample_rate = await _run_on_mlx_thread(_generate_sync)
|
||||||
|
|
||||||
return audio, sample_rate
|
return audio, sample_rate
|
||||||
|
|
||||||
@@ -293,7 +308,7 @@ class MLXSTTBackend:
|
|||||||
return
|
return
|
||||||
|
|
||||||
# Run blocking load in thread pool
|
# Run blocking load in thread pool
|
||||||
await asyncio.to_thread(self._load_model_sync, model_size)
|
await _run_on_mlx_thread(self._load_model_sync, model_size)
|
||||||
|
|
||||||
# Alias for compatibility
|
# Alias for compatibility
|
||||||
load_model = load_model_async
|
load_model = load_model_async
|
||||||
@@ -364,4 +379,4 @@ class MLXSTTBackend:
|
|||||||
return str(result).strip()
|
return str(result).strip()
|
||||||
|
|
||||||
# Run blocking transcription in thread pool
|
# Run blocking transcription in thread pool
|
||||||
return await asyncio.to_thread(_transcribe_sync)
|
return await _run_on_mlx_thread(_transcribe_sync)
|
||||||
|
|||||||
Reference in New Issue
Block a user