From 57ef6636bced0421ba32c84f3244a5c2e7578d9a Mon Sep 17 00:00:00 2001 From: Ron David Ben Ishay Date: Tue, 4 Aug 2026 14:40:32 +0400 Subject: [PATCH] fix(mlx): fold reload+inference into one MLX-worker submission CodeRabbit follow-up on the previous round: the asyncio.Lock closed the gap for callers going through generate()/transcribe() consistently, but unload_model() itself never touched that lock and could still land between the load future resolving and the generate/transcribe future being submitted - two separate _run_on_mlx_thread calls, so a real gap existed at the Python level even though the executor is single-worker. Fix: generate() and transcribe() each now submit exactly ONE callable to the MLX executor - a self-healing reload-if-needed-then-infer function - instead of a load submission followed by a separate infer submission. This removes the gap structurally: nothing can observe an intermediate state because there is no intermediate state exposed across an await boundary. unload_model() now also submits unconditionally (loaded-check moved inside _unload_model_sync, which runs atomically with the teardown) rather than racing its own Python-level self.model check against a load in flight. Verified: same-size generate, size-switch generate, unload route, and a post-unload generate all complete clean on the live launchd server. --- backend/backends/mlx_backend.py | 54 ++++++++++++++++++++++++--------- 1 file changed, 40 insertions(+), 14 deletions(-) diff --git a/backend/backends/mlx_backend.py b/backend/backends/mlx_backend.py index 0f02759a..e7b09dd8 100644 --- a/backend/backends/mlx_backend.py +++ b/backend/backends/mlx_backend.py @@ -141,8 +141,13 @@ class MLXTTSBackend: (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() + # Submit unconditionally rather than checking self.model here first: + # that check would race against a load already queued on the worker + # thread (this call could see None, skip, and leave a model that + # finishes loading a moment later still resident). The loaded check + # belongs inside _unload_model_sync, where it runs atomically with + # the teardown itself. + _mlx_executor.submit(self._unload_model_sync).result() def _unload_model_sync(self): if self.model is not None: @@ -296,12 +301,23 @@ class MLXTTSBackend: return audio, sample_rate - # 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). + def _reload_and_generate_sync(): + """Ensure the configured model is loaded, then generate — as ONE + MLX-worker submission. Two separate submissions (load, then + generate) leave a gap after the load future resolves and before + the generate future is submitted; an unload_model() call from + another thread could land in that gap and tear down the model + this call is about to use. Folding both into one callable closes + the gap: the executor's own FIFO ordering is the only guarantee + this needs, and the reload check here is self-healing even if an + unload happened to run just before this callable started. + """ + if self.model is None or self._current_model_size != self.model_size: + self._reload_sync(self.model_size) + return _generate_sync() + async with self._op_lock: - await self.load_model_async(None) - audio, sample_rate = await _run_on_mlx_thread(_generate_sync) + audio, sample_rate = await _run_on_mlx_thread(_reload_and_generate_sync) return audio, sample_rate @@ -361,10 +377,11 @@ class MLXSTTBackend: def unload_model(self): """Unload the model to free memory. - Safe to call from any thread — see MLXTTSBackend.unload_model for why. + Safe to call from any thread — see MLXTTSBackend.unload_model for why, + including why this submits unconditionally instead of checking + self.model first. """ - if self.model is not None: - _mlx_executor.submit(self._unload_model_sync).result() + _mlx_executor.submit(self._unload_model_sync).result() def _unload_model_sync(self): if self.model is not None: @@ -389,6 +406,8 @@ class MLXSTTBackend: Returns: Transcribed text """ + resolved_size = model_size if model_size is not None else self.model_size + def _transcribe_sync(): """Run synchronous transcription in thread pool.""" # MLX Whisper transcription using generate method @@ -412,8 +431,15 @@ class MLXSTTBackend: else: return str(result).strip() - # Hold the op lock across load + inference so a concurrent request - # for a different model_size can't swap self.model in between. + def _reload_and_transcribe_sync(): + """Ensure the requested model is loaded, then transcribe — as ONE + MLX-worker submission. See MLXTTSBackend._reload_and_generate_sync + for why this needs to be a single callable rather than a separate + load-then-transcribe pair. + """ + if self.model is None or self.model_size != resolved_size: + self._load_model_sync(resolved_size) + return _transcribe_sync() + async with self._op_lock: - await self.load_model_async(model_size) - return await _run_on_mlx_thread(_transcribe_sync) + return await _run_on_mlx_thread(_reload_and_transcribe_sync)