mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 17:15:19 -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 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):
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
Reference in New Issue
Block a user