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.
This commit is contained in:
jamiepine
2026-10-04 00:01:47 +00:00
committed by capy-ai-staging[bot]
parent cf984885e9
commit dd8ab5cb20
2 changed files with 41 additions and 19 deletions
+14 -9
View File
@@ -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):