mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 09:35:16 -07:00
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:
committed by
capy-ai-staging[bot]
parent
fa8db820d7
commit
57ef6636bc
@@ -141,7 +141,12 @@ 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:
|
||||
# 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):
|
||||
@@ -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,9 +377,10 @@ 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()
|
||||
|
||||
def _unload_model_sync(self):
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user