Compare commits

...
Author SHA1 Message Date
Jamie Pine bcab47f25a fix(offline): remove inference-path HF_HUB_OFFLINE guards
0.4.3 wrapped every inference body (`generate`, `transcribe`,
`create_voice_clone_prompt`) with `force_offline_if_cached(True, …)` to
prevent lazy HF lookups from hanging when the network drops
mid-inference (#462). That trade broke online users: the guard flips
`huggingface_hub.constants.HF_HUB_OFFLINE` globally, so any legitimate
metadata call the library makes during generation (e.g. revision
resolution via `HfApi().model_info`) now raises:

    Cannot reach https://huggingface.co/api/models/Qwen/Qwen3-TTS-…:
    offline mode is enabled.

Hit by multiple users on 0.4.3 within hours of release. The offline
blast radius is much larger than the original hang it fixed.

This reverts the inference-path guards. Load-path guards stay — those
worked fine in 0.4.2 and aren't the source of the regression. The
`force_offline_if_cached` helper itself is unchanged; tests still pass.

The #462 hang (network dropping mid-inference) remains unaddressed by
this commit and will need a targeted fix that doesn't flip a global
flag — most likely per-call timeouts or library-specific
`local_files_only` arguments, not a process-wide env mutation.
2026-04-21 04:18:53 -07:00
3 changed files with 80 additions and 96 deletions
+8 -15
View File
@@ -193,8 +193,6 @@ 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,14 +218,12 @@ 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 # Inference runs with the process's default HF_HUB_OFFLINE
# lazy tokenizer/config lookups inside mlx_audio don't hang # state. Forcing offline here (previously used to avoid lazy
# when the user is disconnected (issue #462). # mlx_audio lookups hanging when the network drops mid-inference,
with force_offline_if_cached(True, model_name): # issue #462) regressed online users because libraries make
# Check if model supports voice cloning via generate method # legitimate metadata calls during generation.
# MLX API may support ref_audio parameter directly
try: try:
# Try with voice cloning parameters if supported
if ref_audio: if ref_audio:
# Check if generate accepts ref_audio parameter # Check if generate accepts ref_audio parameter
import inspect import inspect
@@ -347,8 +343,6 @@ 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
@@ -357,10 +351,9 @@ class MLXSTTBackend:
if language: if language:
decode_options["language"] = language decode_options["language"] = language
# Model is loaded → weights are on disk. Force offline so # Inference runs with the process's default HF_HUB_OFFLINE
# lazy tokenizer/config lookups don't hang when the user is # state — see the comment in MLXTTSBackend.generate for the
# disconnected (issue #462). # regression this revert fixes (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
+10 -18
View File
@@ -172,14 +172,12 @@ 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 # Inference runs with the process's default HF_HUB_OFFLINE
# lazy tokenizer/config lookups inside qwen_tts don't hang # state. Forcing offline here (issue #462) regressed online
# when the user is disconnected (issue #462). # users whose libraries issue legitimate metadata lookups
with force_offline_if_cached(True, model_name): # during voice-prompt creation.
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,
@@ -227,18 +225,14 @@ 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)
# Model is loaded → weights are on disk. Force offline so # See _create_prompt_sync comment — inference runs with the
# lazy tokenizer/config lookups inside qwen_tts don't hang # process's default HF_HUB_OFFLINE state (issue #462).
# 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,
@@ -342,17 +336,15 @@ 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 # Inference runs with the process's default HF_HUB_OFFLINE
# `get_decoder_prompt_ids` and any lazy tokenizer lookups # state — forcing offline here (issue #462) broke online users
# don't hang when the user is disconnected (issue #462). # whose `get_decoder_prompt_ids` / tokenizer calls issue
with force_offline_if_cached(True, progress_model_name): # legitimate metadata lookups.
# Process audio # Process audio
inputs = self.processor( inputs = self.processor(
audio, audio,
@@ -186,7 +186,6 @@ 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:
@@ -206,10 +205,10 @@ class QwenCustomVoiceBackend:
if instruct: if instruct:
kwargs["instruct"] = instruct kwargs["instruct"] = instruct
# Model is loaded → weights are on disk. Force offline so # Inference runs with the process's default HF_HUB_OFFLINE
# lazy tokenizer/config lookups inside qwen_tts don't hang # state. Forcing offline here (issue #462) regressed online
# when the user is disconnected (issue #462). # users whose libraries issue legitimate metadata lookups
with force_offline_if_cached(True, model_name): # during generation.
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