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.
"""
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.