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.
This commit is contained in:
Ron David Ben Ishay
2026-10-04 00:01:47 +00:00
committed by capy-ai-staging[bot]
parent fa8db820d7
commit 57ef6636bc
+38 -12
View File
@@ -141,7 +141,12 @@ class MLXTTSBackend:
(e.g. from _reload_sync) — use _unload_model_sync directly there, or (e.g. from _reload_sync) — use _unload_model_sync directly there, or
this would deadlock the single-worker executor waiting on itself. this would deadlock the single-worker executor waiting on itself.
""" """
if self.model is not None: # 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() _mlx_executor.submit(self._unload_model_sync).result()
def _unload_model_sync(self): def _unload_model_sync(self):
@@ -296,12 +301,23 @@ class MLXTTSBackend:
return audio, sample_rate return audio, sample_rate
# Hold the op lock across load + inference so a concurrent request def _reload_and_generate_sync():
# for a different model_size can't swap self.model in between (the """Ensure the configured model is loaded, then generate — as ONE
# model is read inside _generate_sync via closure, after this point). 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: async with self._op_lock:
await self.load_model_async(None) audio, sample_rate = await _run_on_mlx_thread(_reload_and_generate_sync)
audio, sample_rate = await _run_on_mlx_thread(_generate_sync)
return audio, sample_rate return audio, sample_rate
@@ -361,9 +377,10 @@ class MLXSTTBackend:
def unload_model(self): 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. 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): def _unload_model_sync(self):
@@ -389,6 +406,8 @@ class MLXSTTBackend:
Returns: Returns:
Transcribed text Transcribed text
""" """
resolved_size = model_size if model_size is not None else self.model_size
def _transcribe_sync(): def _transcribe_sync():
"""Run synchronous transcription in thread pool.""" """Run synchronous transcription in thread pool."""
# MLX Whisper transcription using generate method # MLX Whisper transcription using generate method
@@ -412,8 +431,15 @@ class MLXSTTBackend:
else: else:
return str(result).strip() return str(result).strip()
# Hold the op lock across load + inference so a concurrent request def _reload_and_transcribe_sync():
# for a different model_size can't swap self.model in between. """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: async with self._op_lock:
await self.load_model_async(model_size) return await _run_on_mlx_thread(_reload_and_transcribe_sync)
return await _run_on_mlx_thread(_transcribe_sync)