mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 09:05:17 -07:00
fix(offline): guard inference paths with HF_HUB_OFFLINE (#462)
PR #443 wrapped the model *load* path with `force_offline_if_cached` so cached models don't phone home at startup. The context manager restores `HF_HUB_OFFLINE` on exit, which left inference paths (generate, transcribe, voice-prompt creation) unguarded — and `qwen_tts`, `mlx_audio`, and `transformers` perform lazy tokenizer/processor/config lookups during inference. With internet on, those lookups are near-instant and invisible; with internet off, `requests` hangs on DNS or connect until the network returns. This is exactly what users in #462 describe: model shows "Loaded", internet drops, generation "thinks" forever, internet comes back, generation completes. Chatterbox and LuxTTS don't exhibit this because their engine libs resolve everything through already-cached paths at load time. Fix: wrap each inference-sync body with `force_offline_if_cached(True, ...)`. Since inference only runs after a successful load, weights are known to be on disk, so `is_cached=True` is unconditional. Also adds the load-time guard that was missing from `qwen_custom_voice_backend.py` — CustomVoice previously had no offline protection at all. Paths patched: - PyTorchTTSBackend.create_voice_prompt (create_voice_clone_prompt) - PyTorchTTSBackend.generate (generate_voice_clone) - PyTorchSTTBackend.transcribe (Whisper generate + decoder-prompt-ids) - MLXTTSBackend.generate (mlx_audio generate, all branches) - MLXSTTBackend.transcribe (mlx_audio whisper generate) - QwenCustomVoiceBackend._load_model_sync + generate Does not address the secondary `check_model_inputs() missing 'func'` error reported in the same issue — that's a `transformers` 5.x version-skew bug on the install path, separate concern. Fixes #462. Co-Authored-By: Claude Opus 4.7 (1M context) <[email protected]>
This commit is contained in:
co-authored by
Claude Opus 4.7
parent
e3f7cd9d00
commit
f3ed312cf2
@@ -195,6 +195,8 @@ class MLXTTSBackend:
|
|||||||
|
|
||||||
logger.info("Generating audio for text: %s", text)
|
logger.info("Generating audio for text: %s", text)
|
||||||
|
|
||||||
|
model_name = f"qwen-tts-{self._current_model_size}"
|
||||||
|
|
||||||
def _generate_sync():
|
def _generate_sync():
|
||||||
"""Run synchronous generation in thread pool."""
|
"""Run synchronous generation in thread pool."""
|
||||||
# MLX generate() returns a generator yielding GenerationResult objects
|
# MLX generate() returns a generator yielding GenerationResult objects
|
||||||
@@ -220,6 +222,10 @@ class MLXTTSBackend:
|
|||||||
logger.warning("Regenerating without voice prompt.")
|
logger.warning("Regenerating without voice prompt.")
|
||||||
ref_audio = None
|
ref_audio = None
|
||||||
|
|
||||||
|
# Model is loaded → weights are on disk. Force offline so
|
||||||
|
# lazy tokenizer/config lookups inside mlx_audio don't hang
|
||||||
|
# when the user is disconnected (issue #462).
|
||||||
|
with force_offline_if_cached(True, model_name):
|
||||||
# Check if model supports voice cloning via generate method
|
# Check if model supports voice cloning via generate method
|
||||||
# MLX API may support ref_audio parameter directly
|
# MLX API may support ref_audio parameter directly
|
||||||
try:
|
try:
|
||||||
@@ -343,6 +349,8 @@ class MLXSTTBackend:
|
|||||||
"""
|
"""
|
||||||
await self.load_model_async(model_size)
|
await self.load_model_async(model_size)
|
||||||
|
|
||||||
|
progress_model_name = f"whisper-{self.model_size}"
|
||||||
|
|
||||||
def _transcribe_sync():
|
def _transcribe_sync():
|
||||||
"""Run synchronous transcription in thread pool."""
|
"""Run synchronous transcription in thread pool."""
|
||||||
# MLX Whisper transcription using generate method
|
# MLX Whisper transcription using generate method
|
||||||
@@ -351,6 +359,10 @@ class MLXSTTBackend:
|
|||||||
if language:
|
if language:
|
||||||
decode_options["language"] = language
|
decode_options["language"] = language
|
||||||
|
|
||||||
|
# Model is loaded → weights are on disk. Force offline so
|
||||||
|
# lazy tokenizer/config lookups don't hang when the user is
|
||||||
|
# disconnected (issue #462).
|
||||||
|
with force_offline_if_cached(True, progress_model_name):
|
||||||
result = self.model.generate(str(audio_path), **decode_options)
|
result = self.model.generate(str(audio_path), **decode_options)
|
||||||
|
|
||||||
# Extract text from result
|
# Extract text from result
|
||||||
|
|||||||
@@ -172,8 +172,14 @@ class PyTorchTTSBackend:
|
|||||||
# This shouldn't happen in practice, but handle it
|
# This shouldn't happen in practice, but handle it
|
||||||
return {"prompt": cached_prompt}, True
|
return {"prompt": cached_prompt}, True
|
||||||
|
|
||||||
|
model_name = f"qwen-tts-{self._current_model_size}"
|
||||||
|
|
||||||
def _create_prompt_sync():
|
def _create_prompt_sync():
|
||||||
"""Run synchronous voice prompt creation in thread pool."""
|
"""Run synchronous voice prompt creation in thread pool."""
|
||||||
|
# Model is loaded → weights are on disk. Force offline so
|
||||||
|
# lazy tokenizer/config lookups inside qwen_tts don't hang
|
||||||
|
# when the user is disconnected (issue #462).
|
||||||
|
with force_offline_if_cached(True, model_name):
|
||||||
return self.model.create_voice_clone_prompt(
|
return self.model.create_voice_clone_prompt(
|
||||||
ref_audio=str(audio_path),
|
ref_audio=str(audio_path),
|
||||||
ref_text=reference_text,
|
ref_text=reference_text,
|
||||||
@@ -221,13 +227,18 @@ class PyTorchTTSBackend:
|
|||||||
# Load model
|
# Load model
|
||||||
await self.load_model_async(None)
|
await self.load_model_async(None)
|
||||||
|
|
||||||
|
model_name = f"qwen-tts-{self._current_model_size}"
|
||||||
|
|
||||||
def _generate_sync():
|
def _generate_sync():
|
||||||
"""Run synchronous generation in thread pool."""
|
"""Run synchronous generation in thread pool."""
|
||||||
# Set seed if provided
|
# Set seed if provided
|
||||||
if seed is not None:
|
if seed is not None:
|
||||||
manual_seed(seed, self.device)
|
manual_seed(seed, self.device)
|
||||||
|
|
||||||
# Generate audio - this is the blocking operation
|
# Model is loaded → weights are on disk. Force offline so
|
||||||
|
# lazy tokenizer/config lookups inside qwen_tts don't hang
|
||||||
|
# when the user is disconnected (issue #462).
|
||||||
|
with force_offline_if_cached(True, model_name):
|
||||||
wavs, sample_rate = self.model.generate_voice_clone(
|
wavs, sample_rate = self.model.generate_voice_clone(
|
||||||
text=text,
|
text=text,
|
||||||
voice_clone_prompt=voice_prompt,
|
voice_clone_prompt=voice_prompt,
|
||||||
@@ -331,11 +342,17 @@ class PyTorchSTTBackend:
|
|||||||
"""
|
"""
|
||||||
await self.load_model_async(model_size)
|
await self.load_model_async(model_size)
|
||||||
|
|
||||||
|
progress_model_name = f"whisper-{self.model_size}"
|
||||||
|
|
||||||
def _transcribe_sync():
|
def _transcribe_sync():
|
||||||
"""Run synchronous transcription in thread pool."""
|
"""Run synchronous transcription in thread pool."""
|
||||||
# Load audio
|
# Load audio
|
||||||
audio, sr = load_audio(audio_path, sample_rate=16000)
|
audio, sr = load_audio(audio_path, sample_rate=16000)
|
||||||
|
|
||||||
|
# Model is loaded → weights are on disk. Force offline so
|
||||||
|
# `get_decoder_prompt_ids` and any lazy tokenizer lookups
|
||||||
|
# don't hang when the user is disconnected (issue #462).
|
||||||
|
with force_offline_if_cached(True, progress_model_name):
|
||||||
# Process audio
|
# Process audio
|
||||||
inputs = self.processor(
|
inputs = self.processor(
|
||||||
audio,
|
audio,
|
||||||
|
|||||||
@@ -28,6 +28,7 @@ from .base import (
|
|||||||
combine_voice_prompts as _combine_voice_prompts,
|
combine_voice_prompts as _combine_voice_prompts,
|
||||||
model_load_progress,
|
model_load_progress,
|
||||||
)
|
)
|
||||||
|
from ..utils.hf_offline_patch import force_offline_if_cached
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -104,6 +105,7 @@ class QwenCustomVoiceBackend:
|
|||||||
model_path = self._get_model_path(model_size)
|
model_path = self._get_model_path(model_size)
|
||||||
logger.info("Loading Qwen CustomVoice %s on %s...", model_size, self.device)
|
logger.info("Loading Qwen CustomVoice %s on %s...", model_size, self.device)
|
||||||
|
|
||||||
|
with force_offline_if_cached(is_cached, model_name):
|
||||||
if self.device == "cpu":
|
if self.device == "cpu":
|
||||||
self.model = Qwen3TTSModel.from_pretrained(
|
self.model = Qwen3TTSModel.from_pretrained(
|
||||||
model_path,
|
model_path,
|
||||||
@@ -184,6 +186,7 @@ class QwenCustomVoiceBackend:
|
|||||||
await self.load_model_async(None)
|
await self.load_model_async(None)
|
||||||
|
|
||||||
speaker = voice_prompt.get("preset_voice_id") or QWEN_CV_DEFAULT_SPEAKER
|
speaker = voice_prompt.get("preset_voice_id") or QWEN_CV_DEFAULT_SPEAKER
|
||||||
|
model_name = f"qwen-custom-voice-{self._current_model_size}"
|
||||||
|
|
||||||
def _generate_sync():
|
def _generate_sync():
|
||||||
if seed is not None:
|
if seed is not None:
|
||||||
@@ -203,6 +206,10 @@ class QwenCustomVoiceBackend:
|
|||||||
if instruct:
|
if instruct:
|
||||||
kwargs["instruct"] = instruct
|
kwargs["instruct"] = instruct
|
||||||
|
|
||||||
|
# Model is loaded → weights are on disk. Force offline so
|
||||||
|
# lazy tokenizer/config lookups inside qwen_tts don't hang
|
||||||
|
# when the user is disconnected (issue #462).
|
||||||
|
with force_offline_if_cached(True, model_name):
|
||||||
wavs, sample_rate = self.model.generate_custom_voice(**kwargs)
|
wavs, sample_rate = self.model.generate_custom_voice(**kwargs)
|
||||||
return wavs[0], sample_rate
|
return wavs[0], sample_rate
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user