mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 09:35:16 -07:00
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:
@@ -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,7 +154,6 @@ 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)
|
||||||
|
|
||||||
@@ -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
|
||||||
|
# the model's default voice would silently break the clone the
|
||||||
|
# user asked for.
|
||||||
if ref_audio:
|
if ref_audio:
|
||||||
# Check if generate accepts ref_audio parameter
|
|
||||||
import inspect
|
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(
|
||||||
|
"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):
|
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))
|
audio_chunks.append(np.array(result.audio))
|
||||||
sample_rate = result.sample_rate
|
sample_rate = result.sample_rate
|
||||||
else:
|
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)
|
|
||||||
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])
|
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
|
||||||
|
|
||||||
# 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,11 +365,10 @@ 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
|
||||||
|
|||||||
Reference in New Issue
Block a user