mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 01:25:18 -07:00
fix(mlx): release the model binding inside the sync bodies so a failing op's traceback cannot pin it
Review follow-up: the post-op drain ran in a finally, but on the error path the exception's traceback kept the _generate_sync/_transcribe_sync frame (and its model local) alive, so mx.clear_cache() had nothing to return. Each sync body now drops its model (and tokenizer) binding in its own finally before the exception propagates.
This commit is contained in:
committed by
capy-ai-staging[bot]
parent
fb67ad436c
commit
63899fd865
@@ -241,70 +241,76 @@ class MLXTTSBackend:
|
|||||||
# thread mid-generation cannot turn a later self.model read into
|
# thread mid-generation cannot turn a later self.model read into
|
||||||
# None (the fallback path below runs seconds into a request).
|
# None (the fallback path below runs seconds into a request).
|
||||||
model = self.model
|
model = self.model
|
||||||
# MLX generate() returns a generator yielding GenerationResult objects
|
|
||||||
audio_chunks = []
|
|
||||||
sample_rate = 24000
|
|
||||||
lang = LANGUAGE_CODE_TO_NAME.get(language, "auto")
|
|
||||||
|
|
||||||
# Set seed if provided (MLX uses numpy random)
|
|
||||||
if seed is not None:
|
|
||||||
import mlx.core as mx
|
|
||||||
|
|
||||||
np.random.seed(seed)
|
|
||||||
mx.random.seed(seed)
|
|
||||||
|
|
||||||
# Extract voice prompt info
|
|
||||||
ref_audio = voice_prompt.get("ref_audio") or voice_prompt.get("ref_audio_path")
|
|
||||||
ref_text = voice_prompt.get("ref_text", "")
|
|
||||||
|
|
||||||
# Validate that the audio file exists
|
|
||||||
if ref_audio and not Path(ref_audio).exists():
|
|
||||||
logger.warning("Audio file not found: %s", ref_audio)
|
|
||||||
logger.warning("This may be due to a cached voice prompt referencing a deleted temp file.")
|
|
||||||
logger.warning("Regenerating without voice prompt.")
|
|
||||||
ref_audio = None
|
|
||||||
|
|
||||||
# Inference runs with the process's default HF_HUB_OFFLINE
|
|
||||||
# state. Forcing offline here (previously used to avoid lazy
|
|
||||||
# mlx_audio lookups hanging when the network drops mid-inference,
|
|
||||||
# issue #462) regressed online users because libraries make
|
|
||||||
# legitimate metadata calls during generation.
|
|
||||||
try:
|
try:
|
||||||
if ref_audio:
|
# MLX generate() returns a generator yielding GenerationResult objects
|
||||||
# Check if generate accepts ref_audio parameter
|
audio_chunks = []
|
||||||
import inspect
|
sample_rate = 24000
|
||||||
|
lang = LANGUAGE_CODE_TO_NAME.get(language, "auto")
|
||||||
|
|
||||||
sig = inspect.signature(model.generate)
|
# Set seed if provided (MLX uses numpy random)
|
||||||
if "ref_audio" in sig.parameters:
|
if seed is not None:
|
||||||
# Generate with voice cloning
|
import mlx.core as mx
|
||||||
for result in model.generate(text, ref_audio=ref_audio, ref_text=ref_text, lang_code=lang):
|
|
||||||
audio_chunks.append(np.array(result.audio))
|
np.random.seed(seed)
|
||||||
sample_rate = result.sample_rate
|
mx.random.seed(seed)
|
||||||
|
|
||||||
|
# Extract voice prompt info
|
||||||
|
ref_audio = voice_prompt.get("ref_audio") or voice_prompt.get("ref_audio_path")
|
||||||
|
ref_text = voice_prompt.get("ref_text", "")
|
||||||
|
|
||||||
|
# Validate that the audio file exists
|
||||||
|
if ref_audio and not Path(ref_audio).exists():
|
||||||
|
logger.warning("Audio file not found: %s", ref_audio)
|
||||||
|
logger.warning("This may be due to a cached voice prompt referencing a deleted temp file.")
|
||||||
|
logger.warning("Regenerating without voice prompt.")
|
||||||
|
ref_audio = None
|
||||||
|
|
||||||
|
# Inference runs with the process's default HF_HUB_OFFLINE
|
||||||
|
# state. Forcing offline here (previously used to avoid lazy
|
||||||
|
# mlx_audio lookups hanging when the network drops mid-inference,
|
||||||
|
# issue #462) regressed online users because libraries make
|
||||||
|
# legitimate metadata calls during generation.
|
||||||
|
try:
|
||||||
|
if ref_audio:
|
||||||
|
# Check if generate accepts ref_audio parameter
|
||||||
|
import inspect
|
||||||
|
|
||||||
|
sig = inspect.signature(model.generate)
|
||||||
|
if "ref_audio" in sig.parameters:
|
||||||
|
# Generate with voice cloning
|
||||||
|
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 model.generate(text, lang_code=lang):
|
||||||
|
audio_chunks.append(np.array(result.audio))
|
||||||
|
sample_rate = result.sample_rate
|
||||||
else:
|
else:
|
||||||
# Fallback: generate without voice cloning
|
# No voice prompt, generate normally
|
||||||
for result in 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:
|
except Exception as e:
|
||||||
# No voice prompt, generate normally
|
# If voice cloning fails, try without it
|
||||||
|
logger.warning("Voice cloning failed, generating without voice prompt: %s", e)
|
||||||
for result in 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:
|
|
||||||
# If voice cloning fails, try without it
|
|
||||||
logger.warning("Voice cloning failed, generating without voice prompt: %s", e)
|
|
||||||
for result in model.generate(text, lang_code=lang):
|
|
||||||
audio_chunks.append(np.array(result.audio))
|
|
||||||
sample_rate = result.sample_rate
|
|
||||||
|
|
||||||
# Concatenate all chunks
|
# Concatenate all chunks
|
||||||
if audio_chunks:
|
if audio_chunks:
|
||||||
audio = np.concatenate([np.asarray(chunk, dtype=np.float32) for chunk in audio_chunks])
|
audio = np.concatenate([np.asarray(chunk, dtype=np.float32) for chunk in audio_chunks])
|
||||||
else:
|
else:
|
||||||
# Fallback: empty audio
|
# Fallback: empty audio
|
||||||
audio = np.array([], dtype=np.float32)
|
audio = np.array([], dtype=np.float32)
|
||||||
|
|
||||||
return audio, sample_rate
|
return audio, sample_rate
|
||||||
|
finally:
|
||||||
|
# Drop the local binding here, inside the frame a propagating
|
||||||
|
# traceback would keep alive, so the caller's drain can
|
||||||
|
# actually return the model's buffers to MLX.
|
||||||
|
del model
|
||||||
|
|
||||||
def _reload_and_generate_sync():
|
def _reload_and_generate_sync():
|
||||||
"""Ensure the configured model is loaded, then generate — as ONE
|
"""Ensure the configured model is loaded, then generate — as ONE
|
||||||
@@ -324,9 +330,10 @@ class MLXTTSBackend:
|
|||||||
finally:
|
finally:
|
||||||
# An unload_model() that landed while we were generating only
|
# An unload_model() that landed while we were generating only
|
||||||
# dropped the backend's reference; the model's buffers were
|
# dropped the backend's reference; the model's buffers were
|
||||||
# kept alive by _generate_sync's local and have just been
|
# kept alive by _generate_sync's local, which that function
|
||||||
# returned to MLX's pool (on success or failure). Drain it now,
|
# drops in its own finally (so a propagating traceback cannot
|
||||||
# or they stay resident until the next load/unload cycle.
|
# pin it), and have just been returned to MLX's pool. Drain it
|
||||||
|
# now, or they stay resident until the next load/unload cycle.
|
||||||
if self.model is None:
|
if self.model is None:
|
||||||
empty_mlx_cache()
|
empty_mlx_cache()
|
||||||
|
|
||||||
@@ -423,24 +430,30 @@ class MLXSTTBackend:
|
|||||||
# 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
|
model = self.model
|
||||||
decode_options = {}
|
try:
|
||||||
if language:
|
decode_options = {}
|
||||||
decode_options["language"] = language
|
if language:
|
||||||
|
decode_options["language"] = language
|
||||||
|
|
||||||
# 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 = 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):
|
||||||
return result.strip()
|
return result.strip()
|
||||||
elif isinstance(result, dict):
|
elif isinstance(result, dict):
|
||||||
return result.get("text", "").strip()
|
return result.get("text", "").strip()
|
||||||
elif hasattr(result, "text"):
|
elif hasattr(result, "text"):
|
||||||
return result.text.strip()
|
return result.text.strip()
|
||||||
else:
|
else:
|
||||||
return str(result).strip()
|
return str(result).strip()
|
||||||
|
finally:
|
||||||
|
# Drop the local binding here, inside the frame a propagating
|
||||||
|
# traceback would keep alive, so the caller's drain can
|
||||||
|
# actually return the model's buffers to MLX.
|
||||||
|
del model
|
||||||
|
|
||||||
def _reload_and_transcribe_sync():
|
def _reload_and_transcribe_sync():
|
||||||
"""Ensure the requested model is loaded, then transcribe — as ONE
|
"""Ensure the requested model is loaded, then transcribe — as ONE
|
||||||
|
|||||||
@@ -305,21 +305,27 @@ class MLXQwenLLMBackend:
|
|||||||
from mlx_lm.sample_utils import make_sampler
|
from mlx_lm.sample_utils import make_sampler
|
||||||
|
|
||||||
model, tokenizer = self.model, self.tokenizer
|
model, tokenizer = self.model, self.tokenizer
|
||||||
messages = _build_messages(prompt, system, examples)
|
try:
|
||||||
chat_prompt = tokenizer.apply_chat_template(
|
messages = _build_messages(prompt, system, examples)
|
||||||
messages,
|
chat_prompt = tokenizer.apply_chat_template(
|
||||||
tokenize=False,
|
messages,
|
||||||
add_generation_prompt=True,
|
tokenize=False,
|
||||||
enable_thinking=False,
|
add_generation_prompt=True,
|
||||||
)
|
enable_thinking=False,
|
||||||
|
)
|
||||||
|
|
||||||
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(
|
||||||
model,
|
model,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
prompt=chat_prompt,
|
prompt=chat_prompt,
|
||||||
max_tokens=max_tokens,
|
max_tokens=max_tokens,
|
||||||
sampler=sampler,
|
sampler=sampler,
|
||||||
verbose=False,
|
verbose=False,
|
||||||
)
|
)
|
||||||
return text.strip()
|
return text.strip()
|
||||||
|
finally:
|
||||||
|
# Drop the local binding here, inside the frame a propagating
|
||||||
|
# traceback would keep alive, so the caller's drain can
|
||||||
|
# actually return the model's buffers to MLX.
|
||||||
|
del model, tokenizer
|
||||||
|
|||||||
Reference in New Issue
Block a user