mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 01:25:18 -07:00
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:
committed by
capy-ai-staging[bot]
parent
cf984885e9
commit
dd8ab5cb20
@@ -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):
|
||||
|
||||
@@ -190,6 +190,9 @@ class MLXQwenLLMBackend:
|
||||
self.tokenizer = None
|
||||
self.model_size = model_size
|
||||
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:
|
||||
return self.model is not None
|
||||
@@ -217,10 +220,14 @@ class MLXQwenLLMBackend:
|
||||
# one thread (see mlx_backend._run_on_mlx_thread / issue #699).
|
||||
from .mlx_backend import _run_on_mlx_thread
|
||||
|
||||
if self.model is not None and self._current_model_size != model_size:
|
||||
await _run_on_mlx_thread(self.unload_model)
|
||||
async with self._op_lock:
|
||||
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:
|
||||
from mlx_lm import load as mlx_load
|
||||
@@ -262,12 +269,21 @@ class MLXQwenLLMBackend:
|
||||
model_size: Optional[str] = None,
|
||||
examples: Optional[list[tuple[str, str]]] = None,
|
||||
) -> str:
|
||||
await self.load_model(model_size)
|
||||
from .mlx_backend import _run_on_mlx_thread
|
||||
|
||||
return await _run_on_mlx_thread(
|
||||
self._generate_sync, prompt, system, max_tokens, temperature, examples
|
||||
)
|
||||
resolved_size = model_size if model_size is not None else self.model_size
|
||||
|
||||
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(
|
||||
self,
|
||||
@@ -280,8 +296,9 @@ class MLXQwenLLMBackend:
|
||||
from mlx_lm import generate as mlx_generate
|
||||
from mlx_lm.sample_utils import make_sampler
|
||||
|
||||
model, tokenizer = self.model, self.tokenizer
|
||||
messages = _build_messages(prompt, system, examples)
|
||||
chat_prompt = self.tokenizer.apply_chat_template(
|
||||
chat_prompt = tokenizer.apply_chat_template(
|
||||
messages,
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
@@ -290,8 +307,8 @@ class MLXQwenLLMBackend:
|
||||
|
||||
sampler = make_sampler(temp=temperature, top_p=0.9) if temperature > 0 else None
|
||||
text = mlx_generate(
|
||||
self.model,
|
||||
self.tokenizer,
|
||||
model,
|
||||
tokenizer,
|
||||
prompt=chat_prompt,
|
||||
max_tokens=max_tokens,
|
||||
sampler=sampler,
|
||||
|
||||
Reference in New Issue
Block a user