From fa8db820d7775067d8d87b45bd02034e36153050 Mon Sep 17 00:00:00 2001 From: Ron David Ben Ishay Date: Tue, 4 Aug 2026 14:13:22 +0400 Subject: [PATCH] fix(mlx): address review feedback on the thread-affinity fix Two real follow-on issues found by CodeRabbit on PR #989: 1. MLXTTSBackend.load_model_async called self.unload_model() directly on the event-loop thread before dispatching _load_model_sync to the MLX worker thread - model teardown ran on the wrong OS thread, same class of bug the PR itself fixes. Combined unload+load into one _reload_sync callable submitted as a single MLX-thread operation. 2. Public unload_model() (called synchronously from the /models/unload routes) also ran on the caller thread. Now submits to the MLX executor and blocks on the result, so teardown always happens on the worker thread regardless of caller. Same fix applied to MLXSTTBackend. 3. Both backends cache self.model on the instance, but generate()/ transcribe() only locked their own internal load step - a concurrent request for a different model_size could swap self.model between one request's load and its inference (real race: routes/ generations.py calls load_engine_model() and generate_chunked() as separate awaited steps with a gap between them). Added a per-backend asyncio.Lock held across the full load+inference sequence in both generate() and transcribe(). Note: this closes the race for callers using the backend's own public methods consistently; the wider route-level orchestration race (load_engine_model + generate_chunked as two separate calls) is a follow-up outside this file's scope. Verified: same-size and cross-size-switch generations both complete clean after the patch; /models/{name}/unload route returns 200 without deadlocking the (still-responsive) server. --- backend/backends/mlx_backend.py | 71 +++++++++++++++++++++++++-------- 1 file changed, 54 insertions(+), 17 deletions(-) diff --git a/backend/backends/mlx_backend.py b/backend/backends/mlx_backend.py index ae87cf6b..0f02759a 100644 --- a/backend/backends/mlx_backend.py +++ b/backend/backends/mlx_backend.py @@ -44,6 +44,10 @@ class MLXTTSBackend: self.model = None self.model_size = model_size self._current_model_size = None + # Guards the whole load-then-use sequence in generate()/create_voice_prompt() + # so a concurrent request for a different model_size can't swap self.model + # out from under an in-flight request between its load and its inference. + self._op_lock = asyncio.Lock() def is_loaded(self) -> bool: """Check if model is loaded.""" @@ -92,16 +96,23 @@ class MLXTTSBackend: if self.model is not None and self._current_model_size == model_size: return - # Unload existing model if different size requested - if self.model is not None and self._current_model_size != model_size: - self.unload_model() - - # Run blocking load in thread pool - await _run_on_mlx_thread(self._load_model_sync, model_size) + # Unload (if needed) and load as ONE callable on the MLX worker thread. + # Doing this as two separate _run_on_mlx_thread calls would run the + # unload on whichever thread issues the second call — usually still + # correct, but a caller-side await gap between them would let another + # coroutine slip a conflicting load in between. One callable removes + # the gap. + await _run_on_mlx_thread(self._reload_sync, model_size) # Alias for compatibility load_model = load_model_async + def _reload_sync(self, model_size: str): + """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_sync() + self._load_model_sync(model_size) + def _load_model_sync(self, model_size: str): """Synchronous model loading.""" model_path = self._get_model_path(model_size) @@ -120,7 +131,20 @@ class MLXTTSBackend: logger.info("MLX TTS model %s loaded successfully", model_size) def unload_model(self): - """Unload the model to free memory.""" + """Unload the model to free memory. + + Safe to call from any thread (e.g. the FastAPI event loop, from the + /models/unload routes): the actual teardown is submitted to the MLX + worker thread and awaited synchronously here, so Metal resources are + always released on the same OS thread that created them. Do not call + this from within a callable already running ON the MLX worker thread + (e.g. from _reload_sync) — use _unload_model_sync directly there, or + this would deadlock the single-worker executor waiting on itself. + """ + if self.model is not None: + _mlx_executor.submit(self._unload_model_sync).result() + + def _unload_model_sync(self): if self.model is not None: del self.model self.model = None @@ -147,7 +171,8 @@ class MLXTTSBackend: Returns: Tuple of (voice_prompt_dict, was_cached) """ - await self.load_model_async(None) + async with self._op_lock: + await self.load_model_async(None) # Check cache if enabled if use_cache: @@ -202,8 +227,6 @@ class MLXTTSBackend: Returns: Tuple of (audio_array, sample_rate) """ - await self.load_model_async(None) - logger.info("Generating audio for text: %s", text) def _generate_sync(): @@ -273,8 +296,12 @@ class MLXTTSBackend: return audio, sample_rate - # Run blocking inference in thread pool - audio, sample_rate = await _run_on_mlx_thread(_generate_sync) + # Hold the op lock across load + inference so a concurrent request + # for a different model_size can't swap self.model in between (the + # model is read inside _generate_sync via closure, after this point). + async with self._op_lock: + await self.load_model_async(None) + audio, sample_rate = await _run_on_mlx_thread(_generate_sync) return audio, sample_rate @@ -285,6 +312,8 @@ class MLXSTTBackend: def __init__(self, model_size: str = "base"): self.model = None self.model_size = model_size + # See MLXTTSBackend._op_lock — same reason. + self._op_lock = asyncio.Lock() def is_loaded(self) -> bool: """Check if model is loaded.""" @@ -330,7 +359,14 @@ class MLXSTTBackend: logger.info("MLX Whisper model %s loaded successfully", model_size) def unload_model(self): - """Unload the model to free memory.""" + """Unload the model to free memory. + + Safe to call from any thread — see MLXTTSBackend.unload_model for why. + """ + if self.model is not None: + _mlx_executor.submit(self._unload_model_sync).result() + + def _unload_model_sync(self): if self.model is not None: del self.model self.model = None @@ -353,8 +389,6 @@ class MLXSTTBackend: Returns: Transcribed text """ - await self.load_model_async(model_size) - def _transcribe_sync(): """Run synchronous transcription in thread pool.""" # MLX Whisper transcription using generate method @@ -378,5 +412,8 @@ class MLXSTTBackend: else: return str(result).strip() - # Run blocking transcription in thread pool - return await _run_on_mlx_thread(_transcribe_sync) + # Hold the op lock across load + inference so a concurrent request + # for a different model_size can't swap self.model in between. + async with self._op_lock: + await self.load_model_async(model_size) + return await _run_on_mlx_thread(_transcribe_sync)