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 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):
|
||||
|
||||
Reference in New Issue
Block a user