From dd8ab5cb20768949d00842e87025f326c59540df Mon Sep 17 00:00:00 2001 From: jamiepine <32987599+jamiepine@users.noreply.github.com> Date: Sat, 3 Oct 2026 20:18:52 +0000 Subject: [PATCH] fix(mlx): fold Qwen3 LLM load+generate into one worker submission; bind model locally in generate closures Review follow-ups on the thread-affinity PR: MLXQwenLLMBackend.generate awaited load_model as one worker submission and then submitted _generate_sync as a second, leaving the same load/generate gap the TTS and STT paths close (a concurrent load_model for another size or an unload could swap or null out self.model/self.tokenizer in between). Give the LLM backend the same _op_lock + single _reload_and_generate_sync shape. The TTS/STT/LLM sync closures now bind the model to a local once, so an inline unload_model() from the event loop mid-generation cannot turn a later self.model read (the voice-clone fallback path in particular) into an AttributeError. --- backend/backends/mlx_backend.py | 23 ++++++++++------- backend/backends/qwen_llm_backend.py | 37 ++++++++++++++++++++-------- 2 files changed, 41 insertions(+), 19 deletions(-) diff --git a/backend/backends/mlx_backend.py b/backend/backends/mlx_backend.py index 2457eeeb..ce3cdc31 100644 --- a/backend/backends/mlx_backend.py +++ b/backend/backends/mlx_backend.py @@ -138,9 +138,9 @@ class MLXTTSBackend: through its global allocator, no stream needed), and routing it through the single worker would block the caller — usually the FastAPI event loop, via /models/unload — until any in-flight - generation on that worker finishes. A generation still running keeps - its own reference to the model, so it completes normally and the - next generate() reloads via _reload_and_generate_sync. + generation on that worker finishes. A generation still running bound + the model to a local before it started, so it completes normally and + the next generate() reloads via _reload_and_generate_sync. """ self._unload_model_sync() @@ -231,6 +231,10 @@ class MLXTTSBackend: def _generate_sync(): """Run synchronous generation in thread pool.""" + # Bind the model once so an inline unload_model() from another + # thread mid-generation cannot turn a later self.model read into + # None (the fallback path below runs seconds into a request). + model = self.model # MLX generate() returns a generator yielding GenerationResult objects audio_chunks = [] sample_rate = 24000 @@ -264,26 +268,26 @@ class MLXTTSBackend: # Check if generate accepts ref_audio parameter import inspect - sig = inspect.signature(self.model.generate) + sig = inspect.signature(model.generate) if "ref_audio" in sig.parameters: # Generate with voice cloning - for result in self.model.generate(text, ref_audio=ref_audio, ref_text=ref_text, lang_code=lang): + for result in model.generate(text, ref_audio=ref_audio, ref_text=ref_text, lang_code=lang): audio_chunks.append(np.array(result.audio)) sample_rate = result.sample_rate else: # Fallback: generate without voice cloning - for result in self.model.generate(text, lang_code=lang): + for result in model.generate(text, lang_code=lang): audio_chunks.append(np.array(result.audio)) sample_rate = result.sample_rate else: # No voice prompt, generate normally - for result in self.model.generate(text, lang_code=lang): + for result in model.generate(text, lang_code=lang): audio_chunks.append(np.array(result.audio)) sample_rate = result.sample_rate except Exception as e: # If voice cloning fails, try without it logger.warning("Voice cloning failed, generating without voice prompt: %s", e) - for result in self.model.generate(text, lang_code=lang): + for result in model.generate(text, lang_code=lang): audio_chunks.append(np.array(result.audio)) sample_rate = result.sample_rate @@ -402,6 +406,7 @@ class MLXSTTBackend: """Run synchronous transcription in thread pool.""" # MLX Whisper transcription using generate method # The generate method accepts audio path directly + model = self.model decode_options = {} if language: decode_options["language"] = language @@ -409,7 +414,7 @@ class MLXSTTBackend: # Inference runs with the process's default HF_HUB_OFFLINE # state — see the comment in MLXTTSBackend.generate for the # regression this revert fixes (issue #462). - result = self.model.generate(str(audio_path), **decode_options) + result = model.generate(str(audio_path), **decode_options) # Extract text from result if isinstance(result, str): diff --git a/backend/backends/qwen_llm_backend.py b/backend/backends/qwen_llm_backend.py index c2ec70c3..a7b99632 100644 --- a/backend/backends/qwen_llm_backend.py +++ b/backend/backends/qwen_llm_backend.py @@ -190,6 +190,9 @@ class MLXQwenLLMBackend: self.tokenizer = None self.model_size = model_size self._current_model_size: Optional[str] = None + # Same role as MLXTTSBackend._op_lock: keeps two coroutines from + # racing their reload decisions around one load-then-generate. + self._op_lock = asyncio.Lock() def is_loaded(self) -> bool: return self.model is not None @@ -217,10 +220,14 @@ class MLXQwenLLMBackend: # one thread (see mlx_backend._run_on_mlx_thread / issue #699). from .mlx_backend import _run_on_mlx_thread - if self.model is not None and self._current_model_size != model_size: - await _run_on_mlx_thread(self.unload_model) + async with self._op_lock: + await _run_on_mlx_thread(self._reload_sync, model_size) - await _run_on_mlx_thread(self._load_model_sync, model_size) + def _reload_sync(self, model_size: str) -> None: + """Unload a mismatched model and load the requested one, in one MLX-thread op.""" + if self.model is not None and self._current_model_size != model_size: + self.unload_model() + self._load_model_sync(model_size) def _load_model_sync(self, model_size: str) -> None: from mlx_lm import load as mlx_load @@ -262,12 +269,21 @@ class MLXQwenLLMBackend: model_size: Optional[str] = None, examples: Optional[list[tuple[str, str]]] = None, ) -> str: - await self.load_model(model_size) from .mlx_backend import _run_on_mlx_thread - return await _run_on_mlx_thread( - self._generate_sync, prompt, system, max_tokens, temperature, examples - ) + resolved_size = model_size if model_size is not None else self.model_size + + def _reload_and_generate_sync() -> str: + # One worker submission for load + generate, as in + # MLXTTSBackend._reload_and_generate_sync: no gap between the two + # where another request's load_model (different size) or an + # unload can swap or null out self.model / self.tokenizer. + if self.model is None or self._current_model_size != resolved_size: + self._reload_sync(resolved_size) + return self._generate_sync(prompt, system, max_tokens, temperature, examples) + + async with self._op_lock: + return await _run_on_mlx_thread(_reload_and_generate_sync) def _generate_sync( self, @@ -280,8 +296,9 @@ class MLXQwenLLMBackend: from mlx_lm import generate as mlx_generate from mlx_lm.sample_utils import make_sampler + model, tokenizer = self.model, self.tokenizer messages = _build_messages(prompt, system, examples) - chat_prompt = self.tokenizer.apply_chat_template( + chat_prompt = tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True, @@ -290,8 +307,8 @@ class MLXQwenLLMBackend: sampler = make_sampler(temp=temperature, top_p=0.9) if temperature > 0 else None text = mlx_generate( - self.model, - self.tokenizer, + model, + tokenizer, prompt=chat_prompt, max_tokens=max_tokens, sampler=sampler,