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:
jamiepine
2026-10-04 00:25:53 +00:00
committed by capy-ai-staging[bot]
parent fb67ad436c
commit 63899fd865
2 changed files with 108 additions and 89 deletions
+85 -72
View File
@@ -241,70 +241,76 @@ class MLXTTSBackend:
# 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
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:
if ref_audio:
# Check if generate accepts ref_audio parameter
import inspect
# MLX generate() returns a generator yielding GenerationResult objects
audio_chunks = []
sample_rate = 24000
lang = LANGUAGE_CODE_TO_NAME.get(language, "auto")
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
# 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:
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:
# Fallback: generate without voice cloning
# No voice prompt, generate normally
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
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
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
if audio_chunks:
audio = np.concatenate([np.asarray(chunk, dtype=np.float32) for chunk in audio_chunks])
else:
# Fallback: empty audio
audio = np.array([], dtype=np.float32)
# Concatenate all chunks
if audio_chunks:
audio = np.concatenate([np.asarray(chunk, dtype=np.float32) for chunk in audio_chunks])
else:
# Fallback: empty audio
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():
"""Ensure the configured model is loaded, then generate — as ONE
@@ -324,9 +330,10 @@ class MLXTTSBackend:
finally:
# An unload_model() that landed while we were generating only
# dropped the backend's reference; the model's buffers were
# kept alive by _generate_sync's local and have just been
# returned to MLX's pool (on success or failure). Drain it now,
# or they stay resident until the next load/unload cycle.
# kept alive by _generate_sync's local, which that function
# drops in its own finally (so a propagating traceback cannot
# 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:
empty_mlx_cache()
@@ -423,24 +430,30 @@ class MLXSTTBackend:
# MLX Whisper transcription using generate method
# The generate method accepts audio path directly
model = self.model
decode_options = {}
if language:
decode_options["language"] = language
try:
decode_options = {}
if language:
decode_options["language"] = language
# 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 = model.generate(str(audio_path), **decode_options)
# 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 = model.generate(str(audio_path), **decode_options)
# Extract text from result
if isinstance(result, str):
return result.strip()
elif isinstance(result, dict):
return result.get("text", "").strip()
elif hasattr(result, "text"):
return result.text.strip()
else:
return str(result).strip()
# Extract text from result
if isinstance(result, str):
return result.strip()
elif isinstance(result, dict):
return result.get("text", "").strip()
elif hasattr(result, "text"):
return result.text.strip()
else:
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():
"""Ensure the requested model is loaded, then transcribe — as ONE
+23 -17
View File
@@ -305,21 +305,27 @@ class MLXQwenLLMBackend:
from mlx_lm.sample_utils import make_sampler
model, tokenizer = self.model, self.tokenizer
messages = _build_messages(prompt, system, examples)
chat_prompt = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
enable_thinking=False,
)
try:
messages = _build_messages(prompt, system, examples)
chat_prompt = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
enable_thinking=False,
)
sampler = make_sampler(temp=temperature, top_p=0.9) if temperature > 0 else None
text = mlx_generate(
model,
tokenizer,
prompt=chat_prompt,
max_tokens=max_tokens,
sampler=sampler,
verbose=False,
)
return text.strip()
sampler = make_sampler(temp=temperature, top_p=0.9) if temperature > 0 else None
text = mlx_generate(
model,
tokenizer,
prompt=chat_prompt,
max_tokens=max_tokens,
sampler=sampler,
verbose=False,
)
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