From 63899fd865d60df849d164fdf38ce9c3c42e467b Mon Sep 17 00:00:00 2001 From: jamiepine <32987599+jamiepine@users.noreply.github.com> Date: Sun, 4 Oct 2026 00:18:36 +0000 Subject: [PATCH] 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. --- backend/backends/mlx_backend.py | 157 +++++++++++++++------------ backend/backends/qwen_llm_backend.py | 40 ++++--- 2 files changed, 108 insertions(+), 89 deletions(-) diff --git a/backend/backends/mlx_backend.py b/backend/backends/mlx_backend.py index 67c5e56f..0f76d95b 100644 --- a/backend/backends/mlx_backend.py +++ b/backend/backends/mlx_backend.py @@ -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 diff --git a/backend/backends/qwen_llm_backend.py b/backend/backends/qwen_llm_backend.py index 81ca7a9f..518afa6e 100644 --- a/backend/backends/qwen_llm_backend.py +++ b/backend/backends/qwen_llm_backend.py @@ -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