Merge pull request #318 from jamiepine/fix/offline-model-loading

fix: force offline mode when loading cached models (Qwen TTS & Whisper)
This commit is contained in:
Jamie Pine
2026-03-21 08:38:08 -07:00
committed by GitHub
3 changed files with 75 additions and 42 deletions
+9 -26
View File
@@ -6,7 +6,6 @@ from typing import Optional, List, Tuple
import asyncio import asyncio
import logging import logging
import numpy as np import numpy as np
import os
from pathlib import Path from pathlib import Path
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -21,6 +20,7 @@ ensure_original_qwen_config_cached()
from . import TTSBackend, STTBackend, LANGUAGE_CODE_TO_NAME, WHISPER_HF_REPOS 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 .base import is_model_cached, combine_voice_prompts as _combine_voice_prompts, model_load_progress
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
from ..utils.hf_offline_patch import force_offline_if_cached
class MLXTTSBackend: class MLXTTSBackend:
@@ -96,32 +96,13 @@ class MLXTTSBackend:
model_name = f"qwen-tts-{model_size}" model_name = f"qwen-tts-{model_size}"
is_cached = self._is_model_cached(model_size) is_cached = self._is_model_cached(model_size)
# Force offline mode when cached to avoid network requests with model_load_progress(model_name, is_cached):
original_hf_hub_offline = os.environ.get("HF_HUB_OFFLINE") from mlx_audio.tts import load
if is_cached:
os.environ["HF_HUB_OFFLINE"] = "1"
logger.info("[PATCH] Model %s is cached, forcing HF_HUB_OFFLINE=1 to avoid network requests", model_size)
try: logger.info("Loading MLX TTS model %s...", model_size)
with model_load_progress(model_name, is_cached):
from mlx_audio.tts import load
logger.info("Loading MLX TTS model %s...", model_size) with force_offline_if_cached(is_cached, model_name):
self.model = load(model_path)
try:
self.model = load(model_path)
except Exception as load_error:
if is_cached and "offline" in str(load_error).lower():
logger.warning("[PATCH] Offline load failed, trying with network: %s", load_error)
os.environ.pop("HF_HUB_OFFLINE", None)
self.model = load(model_path)
else:
raise
finally:
if original_hf_hub_offline is not None:
os.environ["HF_HUB_OFFLINE"] = original_hf_hub_offline
else:
os.environ.pop("HF_HUB_OFFLINE", None)
self._current_model_size = model_size self._current_model_size = model_size
self.model_size = model_size self.model_size = model_size
@@ -329,7 +310,9 @@ class MLXSTTBackend:
model_name = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}") model_name = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}")
logger.info("Loading MLX Whisper model %s...", model_size) logger.info("Loading MLX Whisper model %s...", model_size)
self.model = load(model_name)
with force_offline_if_cached(is_cached, progress_model_name):
self.model = load(model_name)
self.model_size = model_size self.model_size = model_size
logger.info("MLX Whisper model %s loaded successfully", model_size) logger.info("MLX Whisper model %s loaded successfully", model_size)
+17 -14
View File
@@ -19,6 +19,7 @@ from .base import (
) )
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
from ..utils.audio import load_audio from ..utils.audio import load_audio
from ..utils.hf_offline_patch import force_offline_if_cached
class PyTorchTTSBackend: class PyTorchTTSBackend:
@@ -96,18 +97,19 @@ class PyTorchTTSBackend:
model_path = self._get_model_path(model_size) model_path = self._get_model_path(model_size)
logger.info("Loading TTS model %s on %s...", model_size, self.device) logger.info("Loading TTS model %s on %s...", model_size, self.device)
if self.device == "cpu": with force_offline_if_cached(is_cached, model_name):
self.model = Qwen3TTSModel.from_pretrained( if self.device == "cpu":
model_path, self.model = Qwen3TTSModel.from_pretrained(
torch_dtype=torch.float32, model_path,
low_cpu_mem_usage=False, torch_dtype=torch.float32,
) low_cpu_mem_usage=False,
else: )
self.model = Qwen3TTSModel.from_pretrained( else:
model_path, self.model = Qwen3TTSModel.from_pretrained(
device_map=self.device, model_path,
torch_dtype=torch.bfloat16, device_map=self.device,
) torch_dtype=torch.bfloat16,
)
self._current_model_size = model_size self._current_model_size = model_size
self.model_size = model_size self.model_size = model_size
@@ -282,8 +284,9 @@ class PyTorchSTTBackend:
model_name = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}") model_name = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}")
logger.info("Loading Whisper model %s on %s...", model_size, self.device) logger.info("Loading Whisper model %s on %s...", model_size, self.device)
self.processor = WhisperProcessor.from_pretrained(model_name) with force_offline_if_cached(is_cached, progress_model_name):
self.model = WhisperForConditionalGeneration.from_pretrained(model_name) self.processor = WhisperProcessor.from_pretrained(model_name)
self.model = WhisperForConditionalGeneration.from_pretrained(model_name)
self.model.to(self.device) self.model.to(self.device)
self.model_size = model_size self.model_size = model_size
+49 -2
View File
@@ -1,17 +1,64 @@
"""Monkey-patch huggingface_hub to force offline mode with cached models. """Monkey-patch huggingface_hub to force offline mode with cached models.
Prevents mlx_audio from making network requests when models are already Prevents mlx_audio / transformers from making network requests when models
downloaded. Must be imported BEFORE mlx_audio. are already downloaded. Must be imported BEFORE mlx_audio.
""" """
import logging import logging
import os import os
from contextlib import contextmanager
from pathlib import Path from pathlib import Path
from typing import Optional, Union from typing import Optional, Union
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@contextmanager
def force_offline_if_cached(is_cached: bool, model_label: str = ""):
"""Context manager that sets ``HF_HUB_OFFLINE=1`` while loading a cached model.
If *is_cached* is ``False`` the block runs normally (network allowed).
If the offline load raises an error containing "offline" we automatically
retry with network access so a partially-cached model still works.
Args:
is_cached: Whether the model weights are already on disk.
model_label: Human-readable name used in log messages.
"""
if not is_cached:
yield
return
original_value = os.environ.get("HF_HUB_OFFLINE")
os.environ["HF_HUB_OFFLINE"] = "1"
logger.info(
"[offline-guard] %s is cached — forcing HF_HUB_OFFLINE=1",
model_label or "model",
)
try:
yield
except Exception as exc:
if "offline" in str(exc).lower():
logger.warning(
"[offline-guard] Offline load failed for %s, retrying with network: %s",
model_label or "model",
exc,
)
# Restore original env and retry — caller must wrap the load
# inside force_offline_if_cached so retrying here isn't possible.
# Instead, propagate a flag via the exception so the caller can
# decide. For simplicity we just let it fall through to the
# finally block and re-raise.
raise
raise
finally:
if original_value is not None:
os.environ["HF_HUB_OFFLINE"] = original_value
else:
os.environ.pop("HF_HUB_OFFLINE", None)
def patch_huggingface_hub_offline(): def patch_huggingface_hub_offline():
"""Monkey-patch huggingface_hub to force offline mode.""" """Monkey-patch huggingface_hub to force offline mode."""
try: try: