fix(offline): patch transformers mistral-regex check to survive HF failures (#530)

* fix(offline): patch transformers mistral-regex check to survive HF failures

transformers 4.57.x's `PreTrainedTokenizerBase._patch_mistral_regex` calls
`huggingface_hub.model_info(repo_id)` unconditionally during any non-local
tokenizer load to probe for Mistral-family models. The call raises on
`HF_HUB_OFFLINE=1`, on network outages, and on slow/blocked HF endpoints,
and transformers doesn't catch any of it — the exception bubbles out of
`from_pretrained` and kills the load for unrelated engines (Qwen TTS,
Qwen CustomVoice, TADA, etc.).

0.4.2's load-time `force_offline_if_cached` guard walked straight into
this trap: on cached online users it flipped `HF_HUB_OFFLINE=1` and
converted a healthy load into a hard crash. 0.4.3's inference-path guard
masked it; #524 removed the inference guard in 0.4.4, and users updating
to 0.4.4 started hitting the same error on the load path instead
(#526).

Fix:
- Wrap `_patch_mistral_regex` so any exception from the inner HF
  metadata check is swallowed and the tokenizer is returned unchanged.
  Voicebox never loads Mistral models, so the regex rewrite this check
  gates is a no-op for us; matches the success-path behavior for
  non-Mistral repos (tokenization_utils_base.py:2503).
- Drop the `force_offline_if_cached` wraps from every load path
  (pytorch_backend Qwen + Whisper, qwen_custom_voice_backend,
  mlx_backend Qwen + Whisper). With the mistral patch in place they
  provide zero value and only risk re-introducing the same class of
  bug. Helper and its unit tests stay — still correct for targeted
  future use.
- Add `backend/tests/test_offline_patch.py` covering
  OfflineModeIsEnabled / ConnectionError suppression, success
  pass-through, idempotence, and the missing-method no-op path.

Fixes #526.

* fix(offline): install mistral-regex patch for non-MLX backends

The previous commit left the patch wired only through ``mlx_backend.py``'s
existing import of ``hf_offline_patch``. On Windows/Linux/CUDA users who
never load the MLX backend (everyone who hit #526), the patch module was
never imported, so ``patch_transformers_mistral_regex`` never ran and the
crash persisted.

Hoist the import into ``backends/__init__.py``. Every backend imports from
this package, so the module-level patch install runs before any
``from_pretrained`` call regardless of which engine the user picks.

Caught by CodeRabbit and Cursor Bugbot on #530.
This commit is contained in:
Jamie Pine
2026-04-21 22:01:29 -07:00
committed by GitHub
parent 74e004400f
commit d61e884104
6 changed files with 202 additions and 38 deletions
+7
View File
@@ -5,6 +5,13 @@ Provides a unified interface for MLX and PyTorch backends,
and a model config registry that eliminates per-engine dispatch maps.
"""
# Install HF compatibility patches before any backend imports transformers /
# huggingface_hub. The module runs ``patch_transformers_mistral_regex`` at
# import time, which wraps transformers' tokenizer load against the
# unconditional HuggingFace metadata call that otherwise raises on
# HF_HUB_OFFLINE=1 and on network failures.
from ..utils import hf_offline_patch # noqa: F401
import threading
from dataclasses import dataclass, field
from typing import Protocol, Optional, Tuple, List
+2 -5
View File
@@ -20,7 +20,6 @@ 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 ..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:
@@ -99,8 +98,7 @@ class MLXTTSBackend:
logger.info("Loading MLX TTS model %s...", model_size)
with force_offline_if_cached(is_cached, model_name):
self.model = load(model_path)
self.model = load(model_path)
self._current_model_size = model_size
self.model_size = model_size
@@ -311,8 +309,7 @@ class MLXSTTBackend:
model_name = WHISPER_HF_REPOS.get(model_size, f"openai/whisper-{model_size}")
logger.info("Loading MLX Whisper model %s...", model_size)
with force_offline_if_cached(is_cached, progress_model_name):
self.model = load(model_name)
self.model = load(model_name)
self.model_size = model_size
logger.info("MLX Whisper model %s loaded successfully", model_size)
+16 -19
View File
@@ -21,7 +21,6 @@ from .base import (
)
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
from ..utils.audio import load_audio
from ..utils.hf_offline_patch import force_offline_if_cached
class PyTorchTTSBackend:
@@ -106,21 +105,20 @@ class PyTorchTTSBackend:
from huggingface_hub import constants as hf_constants
tts_cache_dir = hf_constants.HF_HUB_CACHE
with force_offline_if_cached(is_cached, model_name):
if self.device == "cpu":
self.model = Qwen3TTSModel.from_pretrained(
model_path,
cache_dir=tts_cache_dir,
torch_dtype=torch.float32,
low_cpu_mem_usage=False,
)
else:
self.model = Qwen3TTSModel.from_pretrained(
model_path,
cache_dir=tts_cache_dir,
device_map=self.device,
torch_dtype=torch.bfloat16,
)
if self.device == "cpu":
self.model = Qwen3TTSModel.from_pretrained(
model_path,
cache_dir=tts_cache_dir,
torch_dtype=torch.float32,
low_cpu_mem_usage=False,
)
else:
self.model = Qwen3TTSModel.from_pretrained(
model_path,
cache_dir=tts_cache_dir,
device_map=self.device,
torch_dtype=torch.bfloat16,
)
self._current_model_size = model_size
self.model_size = model_size
@@ -297,9 +295,8 @@ class PyTorchSTTBackend:
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)
with force_offline_if_cached(is_cached, progress_model_name):
self.processor = WhisperProcessor.from_pretrained(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_size = model_size
+12 -14
View File
@@ -28,7 +28,6 @@ from .base import (
combine_voice_prompts as _combine_voice_prompts,
model_load_progress,
)
from ..utils.hf_offline_patch import force_offline_if_cached
logger = logging.getLogger(__name__)
@@ -105,19 +104,18 @@ class QwenCustomVoiceBackend:
model_path = self._get_model_path(model_size)
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":
self.model = Qwen3TTSModel.from_pretrained(
model_path,
torch_dtype=torch.float32,
low_cpu_mem_usage=False,
)
else:
self.model = Qwen3TTSModel.from_pretrained(
model_path,
device_map=self.device,
torch_dtype=torch.bfloat16,
)
if self.device == "cpu":
self.model = Qwen3TTSModel.from_pretrained(
model_path,
torch_dtype=torch.float32,
low_cpu_mem_usage=False,
)
else:
self.model = Qwen3TTSModel.from_pretrained(
model_path,
device_map=self.device,
torch_dtype=torch.bfloat16,
)
self._current_model_size = model_size
self.model_size = model_size
+113
View File
@@ -0,0 +1,113 @@
"""
Unit tests for ``patch_transformers_mistral_regex``.
Verifies that our wrapper around
``transformers.PreTrainedTokenizerBase._patch_mistral_regex`` catches
exceptions from the unconditional ``huggingface_hub.model_info()`` lookup
and returns the tokenizer unchanged matching the success-path behavior
for non-Mistral repos (transformers 4.57.3, ``tokenization_utils_base.py:2503``).
NOTE: These tests mutate ``transformers.PreTrainedTokenizerBase`` globally;
run serially, not under ``pytest-xdist`` with per-worker process isolation.
"""
import sys
from pathlib import Path
import pytest
sys.path.insert(0, str(Path(__file__).parent.parent))
from huggingface_hub.errors import OfflineModeIsEnabled # noqa: E402
from transformers.tokenization_utils_base import PreTrainedTokenizerBase # noqa: E402
import utils.hf_offline_patch as hf_offline_patch # noqa: E402
@pytest.fixture(autouse=True)
def restore_mistral_regex():
"""Snapshot the current ``_patch_mistral_regex`` and restore after each test."""
saved = PreTrainedTokenizerBase.__dict__.get("_patch_mistral_regex")
saved_flag = hf_offline_patch._mistral_regex_patched
try:
yield
finally:
if saved is not None:
PreTrainedTokenizerBase._patch_mistral_regex = saved
hf_offline_patch._mistral_regex_patched = saved_flag
def _apply_patch():
hf_offline_patch._mistral_regex_patched = False
hf_offline_patch.patch_transformers_mistral_regex()
def test_suppresses_offline_mode_is_enabled(monkeypatch):
_apply_patch()
import huggingface_hub
def raise_offline(*_args, **_kwargs):
raise OfflineModeIsEnabled("offline")
monkeypatch.setattr(huggingface_hub, "model_info", raise_offline)
sentinel = object()
result = PreTrainedTokenizerBase._patch_mistral_regex(
sentinel, "Qwen/Qwen3-TTS-12Hz-1.7B-Base"
)
assert result is sentinel
def test_suppresses_connection_errors(monkeypatch):
_apply_patch()
import huggingface_hub
def raise_connection(*_args, **_kwargs):
raise ConnectionError("network unreachable")
monkeypatch.setattr(huggingface_hub, "model_info", raise_connection)
sentinel = object()
result = PreTrainedTokenizerBase._patch_mistral_regex(
sentinel, "Qwen/Qwen3-TTS-12Hz-1.7B-Base"
)
assert result is sentinel
def test_passthrough_on_success(monkeypatch):
"""When model_info returns non-Mistral tags the original falls through and returns the tokenizer unchanged."""
_apply_patch()
import huggingface_hub
class FakeInfo:
tags = ["model-type:qwen", "language:en"]
monkeypatch.setattr(huggingface_hub, "model_info", lambda *_a, **_kw: FakeInfo())
sentinel = object()
result = PreTrainedTokenizerBase._patch_mistral_regex(
sentinel, "Qwen/Qwen3-TTS-12Hz-1.7B-Base"
)
assert result is sentinel
def test_idempotent():
_apply_patch()
first = PreTrainedTokenizerBase._patch_mistral_regex
hf_offline_patch.patch_transformers_mistral_regex()
second = PreTrainedTokenizerBase._patch_mistral_regex
assert first.__func__ is second.__func__
def test_missing_method_is_noop(monkeypatch):
monkeypatch.delattr(PreTrainedTokenizerBase, "_patch_mistral_regex", raising=False)
hf_offline_patch._mistral_regex_patched = False
hf_offline_patch.patch_transformers_mistral_regex()
assert hf_offline_patch._mistral_regex_patched is False
if __name__ == "__main__":
pytest.main([__file__, "-v"])
+52
View File
@@ -142,6 +142,57 @@ def force_offline_if_cached(is_cached: bool, model_label: str = ""):
_saved_transformers_const = None
_mistral_regex_patched = False
def patch_transformers_mistral_regex():
"""Make transformers' tokenizer load robust to HuggingFace metadata failures.
transformers 4.57.x added ``PreTrainedTokenizerBase._patch_mistral_regex``
which unconditionally calls ``huggingface_hub.model_info(repo_id)`` during
every non-local tokenizer load to check whether the model is a Mistral
variant. That call raises on ``HF_HUB_OFFLINE=1`` and on plain network
failures, killing unrelated loads (Qwen TTS, TADA, etc.).
Voicebox never loads Mistral models, so the rewrite the function would
apply is a no-op for us anyway. Wrap the method so any exception from the
metadata lookup returns the tokenizer unchanged matching the success-path
behavior for non-Mistral repos (transformers 4.57.3,
``tokenization_utils_base.py:2503``).
"""
global _mistral_regex_patched
if _mistral_regex_patched:
return
try:
from transformers.tokenization_utils_base import PreTrainedTokenizerBase
except ImportError:
logger.debug("transformers not available, skipping mistral-regex patch")
return
original = getattr(PreTrainedTokenizerBase, "_patch_mistral_regex", None)
if original is None:
logger.debug(
"transformers has no _patch_mistral_regex attribute, skipping patch",
)
return
def safe_patch_mistral_regex(cls, tokenizer, pretrained_model_name_or_path, *args, **kwargs):
try:
return original(tokenizer, pretrained_model_name_or_path, *args, **kwargs)
except Exception as exc:
logger.debug(
"[mistral-regex-patch] suppressed %s for %r, returning tokenizer as-is",
type(exc).__name__,
pretrained_model_name_or_path,
)
return tokenizer
PreTrainedTokenizerBase._patch_mistral_regex = classmethod(safe_patch_mistral_regex)
_mistral_regex_patched = True
logger.debug("installed _patch_mistral_regex wrapper")
def patch_huggingface_hub_offline():
"""Monkey-patch huggingface_hub to force offline mode."""
try:
@@ -215,4 +266,5 @@ def ensure_original_qwen_config_cached():
if os.environ.get("VOICEBOX_OFFLINE_PATCH", "1") != "0":
patch_huggingface_hub_offline()
patch_transformers_mistral_regex()
ensure_original_qwen_config_cached()