From 11934c2b7d63021a14cc91f4d64619393037231e Mon Sep 17 00:00:00 2001 From: Jamie Pine Date: Sun, 26 Jul 2026 23:17:08 -0700 Subject: [PATCH] fix(mlx): fail generation when voice cloning fails A cloning failure was caught and retried without the voice prompt, so the user got the model's default voice recorded as a successful generation. The error now propagates and the worker records the generation as failed with the real message. Two sibling silent paths raise as well: a model whose generate() lacks ref_audio support, and a generation that produces no audio chunks. --- backend/backends/mlx_backend.py | 93 ++++++++++++++------------------- 1 file changed, 40 insertions(+), 53 deletions(-) diff --git a/backend/backends/mlx_backend.py b/backend/backends/mlx_backend.py index 92ba80ee..323c62d3 100644 --- a/backend/backends/mlx_backend.py +++ b/backend/backends/mlx_backend.py @@ -2,24 +2,24 @@ MLX backend implementation for TTS and STT using mlx-audio. """ -from typing import Optional, List, Tuple import logging -import numpy as np from pathlib import Path +import numpy as np + logger = logging.getLogger(__name__) # PATCH: Import and apply offline patch BEFORE any huggingface_hub usage # This prevents mlx_audio from making network requests when models are cached -from ..utils.hf_offline_patch import patch_huggingface_hub_offline, ensure_original_qwen_config_cached +from ..utils.hf_offline_patch import ensure_original_qwen_config_cached, patch_huggingface_hub_offline patch_huggingface_hub_offline() ensure_original_qwen_config_cached() -from . import TTSBackend, STTBackend, LANGUAGE_CODE_TO_NAME, WHISPER_HF_REPOS -from .base import is_model_cached, combine_voice_prompts as _combine_voice_prompts, model_load_progress -from ..services.mlx_thread import run_on_mlx_thread, clear_mlx_cache -from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt +from ..services.mlx_thread import clear_mlx_cache, run_on_mlx_thread +from ..utils.cache import cache_voice_prompt, get_cache_key, get_cached_voice_prompt +from . import LANGUAGE_CODE_TO_NAME, WHISPER_HF_REPOS +from .base import combine_voice_prompts as _combine_voice_prompts, is_model_cached, model_load_progress class MLXTTSBackend: @@ -63,7 +63,7 @@ class MLXTTSBackend: weight_extensions=(".safetensors", ".bin", ".npz"), ) - def _ensure_loaded_sync(self, model_size: Optional[str]): + def _ensure_loaded_sync(self, model_size: str | None): """Load the model if the requested size isn't already resident. Runs on the MLX worker thread so it stays serialized with generation. @@ -79,7 +79,7 @@ class MLXTTSBackend: self._load_model_sync(model_size) - async def load_model_async(self, model_size: Optional[str] = None): + async def load_model_async(self, model_size: str | None = None): """ Lazy load the MLX TTS model. @@ -126,7 +126,7 @@ class MLXTTSBackend: audio_path: str, reference_text: str, use_cache: bool = True, - ) -> Tuple[dict, bool]: + ) -> tuple[dict, bool]: """ Create voice prompt from reference audio. @@ -154,9 +154,8 @@ class MLXTTSBackend: cached_audio_path = cached_prompt.get("ref_audio") or cached_prompt.get("ref_audio_path") if cached_audio_path and Path(cached_audio_path).exists(): return cached_prompt, True - else: - # Cached file no longer exists, invalidate cache - logger.warning("Cached audio file not found: %s, regenerating prompt", cached_audio_path) + # Cached file no longer exists, invalidate cache + logger.warning("Cached audio file not found: %s, regenerating prompt", cached_audio_path) # MLX voice prompt format - store audio path and text # The model will process this during generation @@ -180,9 +179,9 @@ class MLXTTSBackend: text: str, voice_prompt: dict, language: str = "en", - seed: Optional[int] = None, - instruct: Optional[str] = None, - ) -> Tuple[np.ndarray, int]: + seed: int | None = None, + instruct: str | None = None, + ) -> tuple[np.ndarray, int]: """ Generate audio from text using voice prompt. @@ -228,41 +227,30 @@ class MLXTTSBackend: # 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 + # A cloning failure surfaces as a failed generation; substituting + # the model's default voice would silently break the clone the + # user asked for. + if ref_audio: + import inspect - sig = inspect.signature(self.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): - 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): - 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): - 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) + sig = inspect.signature(self.model.generate) + if "ref_audio" not in sig.parameters: + raise RuntimeError( + "Loaded MLX model does not support voice cloning " + "(generate() has no ref_audio parameter)" + ) + for result in self.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: for result in self.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) + if not audio_chunks: + raise RuntimeError("Model produced no audio") + audio = np.concatenate([np.asarray(chunk, dtype=np.float32) for chunk in audio_chunks]) return audio, sample_rate # Load-if-needed and inference run as one job on the MLX worker so a @@ -291,7 +279,7 @@ class MLXSTTBackend: hf_repo = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}") return is_model_cached(hf_repo, weight_extensions=(".safetensors", ".bin", ".npz")) - def _ensure_loaded_sync(self, model_size: Optional[str]): + def _ensure_loaded_sync(self, model_size: str | None): """Load the model if the requested size isn't already resident. Runs on the MLX worker thread so it stays serialized with transcription. @@ -304,7 +292,7 @@ class MLXSTTBackend: self._load_model_sync(model_size) - async def load_model_async(self, model_size: Optional[str] = None): + async def load_model_async(self, model_size: str | None = None): """ Lazy load the MLX Whisper model. @@ -347,8 +335,8 @@ class MLXSTTBackend: async def transcribe( self, audio_path: str, - language: Optional[str] = None, - model_size: Optional[str] = None, + language: str | None = None, + model_size: str | None = None, ) -> str: """ Transcribe audio to text. @@ -377,12 +365,11 @@ class MLXSTTBackend: # Extract text from result if isinstance(result, str): return result.strip() - elif isinstance(result, dict): + if isinstance(result, dict): return result.get("text", "").strip() - elif hasattr(result, "text"): + if hasattr(result, "text"): return result.text.strip() - else: - return str(result).strip() + return str(result).strip() # Load-if-needed and transcription run as one job on the MLX worker so # a concurrent unload or load can't land between them.