From f135a471ae00ef2e0d32139756b93ad570ee3b30 Mon Sep 17 00:00:00 2001 From: Ron David Ben Ishay Date: Tue, 4 Aug 2026 13:05:22 +0400 Subject: [PATCH] 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). --- backend/backends/mlx_backend.py | 23 +++++++++++++++++++---- 1 file changed, 19 insertions(+), 4 deletions(-) diff --git a/backend/backends/mlx_backend.py b/backend/backends/mlx_backend.py index 9692e59b..ae87cf6b 100644 --- a/backend/backends/mlx_backend.py +++ b/backend/backends/mlx_backend.py @@ -7,9 +7,24 @@ import asyncio import logging import numpy as np from pathlib import Path +from concurrent.futures import ThreadPoolExecutor 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 # 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 @@ -82,7 +97,7 @@ class MLXTTSBackend: self.unload_model() # 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 load_model = load_model_async @@ -259,7 +274,7 @@ class MLXTTSBackend: return audio, sample_rate # 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 @@ -293,7 +308,7 @@ class MLXSTTBackend: return # 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 load_model = load_model_async @@ -364,4 +379,4 @@ class MLXSTTBackend: return str(result).strip() # Run blocking transcription in thread pool - return await asyncio.to_thread(_transcribe_sync) + return await _run_on_mlx_thread(_transcribe_sync)