mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-09-15 12:50:42 -07:00
* 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.
114 lines
3.4 KiB
Python
114 lines
3.4 KiB
Python
"""
|
|
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"])
|