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.
This commit is contained in:
Jamie Pine
2026-07-26 23:17:08 -07:00
parent 9fddf3f799
commit 11934c2b7d
+40 -53
View File
@@ -2,24 +2,24 @@
MLX backend implementation for TTS and STT using mlx-audio. MLX backend implementation for TTS and STT using mlx-audio.
""" """
from typing import Optional, List, Tuple
import logging import logging
import numpy as np
from pathlib import Path from pathlib import Path
import numpy as np
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# PATCH: Import and apply offline patch BEFORE any huggingface_hub usage # PATCH: Import and apply offline patch BEFORE any huggingface_hub usage
# This prevents mlx_audio from making network requests when models are cached # 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() patch_huggingface_hub_offline()
ensure_original_qwen_config_cached() ensure_original_qwen_config_cached()
from . import TTSBackend, STTBackend, LANGUAGE_CODE_TO_NAME, WHISPER_HF_REPOS from ..services.mlx_thread import clear_mlx_cache, run_on_mlx_thread
from .base import is_model_cached, combine_voice_prompts as _combine_voice_prompts, model_load_progress from ..utils.cache import cache_voice_prompt, get_cache_key, get_cached_voice_prompt
from ..services.mlx_thread import run_on_mlx_thread, clear_mlx_cache from . import LANGUAGE_CODE_TO_NAME, WHISPER_HF_REPOS
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt from .base import combine_voice_prompts as _combine_voice_prompts, is_model_cached, model_load_progress
class MLXTTSBackend: class MLXTTSBackend:
@@ -63,7 +63,7 @@ class MLXTTSBackend:
weight_extensions=(".safetensors", ".bin", ".npz"), 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. """Load the model if the requested size isn't already resident.
Runs on the MLX worker thread so it stays serialized with generation. Runs on the MLX worker thread so it stays serialized with generation.
@@ -79,7 +79,7 @@ class MLXTTSBackend:
self._load_model_sync(model_size) 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. Lazy load the MLX TTS model.
@@ -126,7 +126,7 @@ class MLXTTSBackend:
audio_path: str, audio_path: str,
reference_text: str, reference_text: str,
use_cache: bool = True, use_cache: bool = True,
) -> Tuple[dict, bool]: ) -> tuple[dict, bool]:
""" """
Create voice prompt from reference audio. 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") 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(): if cached_audio_path and Path(cached_audio_path).exists():
return cached_prompt, True return cached_prompt, True
else: # Cached file no longer exists, invalidate cache
# Cached file no longer exists, invalidate cache logger.warning("Cached audio file not found: %s, regenerating prompt", cached_audio_path)
logger.warning("Cached audio file not found: %s, regenerating prompt", cached_audio_path)
# MLX voice prompt format - store audio path and text # MLX voice prompt format - store audio path and text
# The model will process this during generation # The model will process this during generation
@@ -180,9 +179,9 @@ class MLXTTSBackend:
text: str, text: str,
voice_prompt: dict, voice_prompt: dict,
language: str = "en", language: str = "en",
seed: Optional[int] = None, seed: int | None = None,
instruct: Optional[str] = None, instruct: str | None = None,
) -> Tuple[np.ndarray, int]: ) -> tuple[np.ndarray, int]:
""" """
Generate audio from text using voice prompt. Generate audio from text using voice prompt.
@@ -228,41 +227,30 @@ class MLXTTSBackend:
# mlx_audio lookups hanging when the network drops mid-inference, # mlx_audio lookups hanging when the network drops mid-inference,
# issue #462) regressed online users because libraries make # issue #462) regressed online users because libraries make
# legitimate metadata calls during generation. # legitimate metadata calls during generation.
try: # A cloning failure surfaces as a failed generation; substituting
if ref_audio: # the model's default voice would silently break the clone the
# Check if generate accepts ref_audio parameter # user asked for.
import inspect if ref_audio:
import inspect
sig = inspect.signature(self.model.generate) sig = inspect.signature(self.model.generate)
if "ref_audio" in sig.parameters: if "ref_audio" not in sig.parameters:
# Generate with voice cloning raise RuntimeError(
for result in self.model.generate(text, ref_audio=ref_audio, ref_text=ref_text, lang_code=lang): "Loaded MLX model does not support voice cloning "
audio_chunks.append(np.array(result.audio)) "(generate() has no ref_audio parameter)"
sample_rate = result.sample_rate )
else: for result in self.model.generate(text, ref_audio=ref_audio, ref_text=ref_text, lang_code=lang):
# Fallback: generate without voice cloning audio_chunks.append(np.array(result.audio))
for result in self.model.generate(text, lang_code=lang): sample_rate = result.sample_rate
audio_chunks.append(np.array(result.audio)) else:
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)
for result in self.model.generate(text, lang_code=lang): for result in self.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
# Concatenate all chunks if not audio_chunks:
if audio_chunks: raise RuntimeError("Model produced no audio")
audio = np.concatenate([np.asarray(chunk, dtype=np.float32) for chunk in audio_chunks])
else:
# Fallback: empty audio
audio = np.array([], dtype=np.float32)
audio = np.concatenate([np.asarray(chunk, dtype=np.float32) for chunk in audio_chunks])
return audio, sample_rate return audio, sample_rate
# Load-if-needed and inference run as one job on the MLX worker so a # 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}") hf_repo = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}")
return is_model_cached(hf_repo, weight_extensions=(".safetensors", ".bin", ".npz")) 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. """Load the model if the requested size isn't already resident.
Runs on the MLX worker thread so it stays serialized with transcription. Runs on the MLX worker thread so it stays serialized with transcription.
@@ -304,7 +292,7 @@ class MLXSTTBackend:
self._load_model_sync(model_size) 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. Lazy load the MLX Whisper model.
@@ -347,8 +335,8 @@ class MLXSTTBackend:
async def transcribe( async def transcribe(
self, self,
audio_path: str, audio_path: str,
language: Optional[str] = None, language: str | None = None,
model_size: Optional[str] = None, model_size: str | None = None,
) -> str: ) -> str:
""" """
Transcribe audio to text. Transcribe audio to text.
@@ -377,12 +365,11 @@ class MLXSTTBackend:
# 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): if isinstance(result, dict):
return result.get("text", "").strip() return result.get("text", "").strip()
elif hasattr(result, "text"): if hasattr(result, "text"):
return result.text.strip() 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 # Load-if-needed and transcription run as one job on the MLX worker so
# a concurrent unload or load can't land between them. # a concurrent unload or load can't land between them.