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 its global allocator, no stream needed), and routing it
through the single worker would block the caller — usually the through the single worker would block the caller — usually the
FastAPI event loop, via /models/unload — until any in-flight FastAPI event loop, via /models/unload — until any in-flight
generation on that worker finishes. A generation still running keeps generation on that worker finishes. A generation still running bound
its own reference to the model, so it completes normally and the the model to a local before it started, so it completes normally and
next generate() reloads via _reload_and_generate_sync. the next generate() reloads via _reload_and_generate_sync.
""" """
self._unload_model_sync() self._unload_model_sync()
@@ -231,6 +231,10 @@ class MLXTTSBackend:
def _generate_sync(): def _generate_sync():
"""Run synchronous generation in thread pool.""" """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 # MLX generate() returns a generator yielding GenerationResult objects
audio_chunks = [] audio_chunks = []
sample_rate = 24000 sample_rate = 24000
@@ -264,26 +268,26 @@ class MLXTTSBackend:
# Check if generate accepts ref_audio parameter # Check if generate accepts ref_audio parameter
import inspect import inspect
sig = inspect.signature(self.model.generate) sig = inspect.signature(model.generate)
if "ref_audio" in sig.parameters: if "ref_audio" in sig.parameters:
# Generate with voice cloning # 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)) audio_chunks.append(np.array(result.audio))
sample_rate = result.sample_rate sample_rate = result.sample_rate
else: else:
# Fallback: generate without voice cloning # 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)) audio_chunks.append(np.array(result.audio))
sample_rate = result.sample_rate sample_rate = result.sample_rate
else: else:
# No voice prompt, generate normally # 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)) audio_chunks.append(np.array(result.audio))
sample_rate = result.sample_rate sample_rate = result.sample_rate
except Exception as e: except Exception as e:
# If voice cloning fails, try without it # If voice cloning fails, try without it
logger.warning("Voice cloning failed, generating without voice prompt: %s", e) 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)) audio_chunks.append(np.array(result.audio))
sample_rate = result.sample_rate sample_rate = result.sample_rate
@@ -402,6 +406,7 @@ class MLXSTTBackend:
"""Run synchronous transcription in thread pool.""" """Run synchronous transcription in thread pool."""
# MLX Whisper transcription using generate method # MLX Whisper transcription using generate method
# The generate method accepts audio path directly # The generate method accepts audio path directly
model = self.model
decode_options = {} decode_options = {}
if language: if language:
decode_options["language"] = language decode_options["language"] = language
@@ -409,7 +414,7 @@ class MLXSTTBackend:
# Inference runs with the process's default HF_HUB_OFFLINE # Inference runs with the process's default HF_HUB_OFFLINE
# state — see the comment in MLXTTSBackend.generate for the # state — see the comment in MLXTTSBackend.generate for the
# regression this revert fixes (issue #462). # 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 # Extract text from result
if isinstance(result, str): if isinstance(result, str):
+27 -10
View File
@@ -190,6 +190,9 @@ class MLXQwenLLMBackend:
self.tokenizer = None self.tokenizer = None
self.model_size = model_size self.model_size = model_size
self._current_model_size: Optional[str] = None self._current_model_size: Optional[str] = None
# Same role as MLXTTSBackend._op_lock: keeps two coroutines from
# racing their reload decisions around one load-then-generate.
self._op_lock = asyncio.Lock()
def is_loaded(self) -> bool: def is_loaded(self) -> bool:
return self.model is not None return self.model is not None
@@ -217,10 +220,14 @@ class MLXQwenLLMBackend:
# one thread (see mlx_backend._run_on_mlx_thread / issue #699). # one thread (see mlx_backend._run_on_mlx_thread / issue #699).
from .mlx_backend import _run_on_mlx_thread from .mlx_backend import _run_on_mlx_thread
if self.model is not None and self._current_model_size != model_size: async with self._op_lock:
await _run_on_mlx_thread(self.unload_model) await _run_on_mlx_thread(self._reload_sync, model_size)
await _run_on_mlx_thread(self._load_model_sync, model_size) def _reload_sync(self, model_size: str) -> None:
"""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()
self._load_model_sync(model_size)
def _load_model_sync(self, model_size: str) -> None: def _load_model_sync(self, model_size: str) -> None:
from mlx_lm import load as mlx_load from mlx_lm import load as mlx_load
@@ -262,12 +269,21 @@ class MLXQwenLLMBackend:
model_size: Optional[str] = None, model_size: Optional[str] = None,
examples: Optional[list[tuple[str, str]]] = None, examples: Optional[list[tuple[str, str]]] = None,
) -> str: ) -> str:
await self.load_model(model_size)
from .mlx_backend import _run_on_mlx_thread from .mlx_backend import _run_on_mlx_thread
return await _run_on_mlx_thread( resolved_size = model_size if model_size is not None else self.model_size
self._generate_sync, prompt, system, max_tokens, temperature, examples
) def _reload_and_generate_sync() -> str:
# One worker submission for load + generate, as in
# MLXTTSBackend._reload_and_generate_sync: no gap between the two
# where another request's load_model (different size) or an
# unload can swap or null out self.model / self.tokenizer.
if self.model is None or self._current_model_size != resolved_size:
self._reload_sync(resolved_size)
return self._generate_sync(prompt, system, max_tokens, temperature, examples)
async with self._op_lock:
return await _run_on_mlx_thread(_reload_and_generate_sync)
def _generate_sync( def _generate_sync(
self, self,
@@ -280,8 +296,9 @@ class MLXQwenLLMBackend:
from mlx_lm import generate as mlx_generate from mlx_lm import generate as mlx_generate
from mlx_lm.sample_utils import make_sampler from mlx_lm.sample_utils import make_sampler
model, tokenizer = self.model, self.tokenizer
messages = _build_messages(prompt, system, examples) messages = _build_messages(prompt, system, examples)
chat_prompt = self.tokenizer.apply_chat_template( chat_prompt = tokenizer.apply_chat_template(
messages, messages,
tokenize=False, tokenize=False,
add_generation_prompt=True, add_generation_prompt=True,
@@ -290,8 +307,8 @@ class MLXQwenLLMBackend:
sampler = make_sampler(temp=temperature, top_p=0.9) if temperature > 0 else None sampler = make_sampler(temp=temperature, top_p=0.9) if temperature > 0 else None
text = mlx_generate( text = mlx_generate(
self.model, model,
self.tokenizer, tokenizer,
prompt=chat_prompt, prompt=chat_prompt,
max_tokens=max_tokens, max_tokens=max_tokens,
sampler=sampler, sampler=sampler,