mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 17:15:19 -07:00
Merge pull request #1130 from jamiepine/prep/pr-1112
Fix infinite HF retry storm when loading a cached model offline
This commit is contained in:
@@ -0,0 +1,55 @@
|
||||
"""
|
||||
Regression test for voicebox #434 (infinite HF retry loop on a cached model).
|
||||
|
||||
``model_load_progress`` is the single context manager every backend's
|
||||
``_load_model_sync`` uses around its ``from_pretrained()`` call. It already
|
||||
receives ``is_cached`` but never forwarded it to ``force_offline_if_cached``,
|
||||
so a fully-cached model still resolved config files against huggingface.co
|
||||
and ate the default 5-retry backoff per file offline. This asserts the guard
|
||||
is actually active for the duration of the ``with`` block when ``is_cached``
|
||||
is ``True``, and restored on exit.
|
||||
|
||||
NOTE: mutates process-global state in ``huggingface_hub.constants`` and
|
||||
``transformers.utils.hub`` (via ``force_offline_if_cached``); run serially,
|
||||
same caveat as ``test_offline_guard.py``.
|
||||
"""
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent.parent))
|
||||
|
||||
from backend.backends.base import model_load_progress
|
||||
|
||||
|
||||
def _hf_const():
|
||||
import huggingface_hub.constants as hf_const
|
||||
|
||||
return hf_const
|
||||
|
||||
|
||||
def _tf_hub():
|
||||
import transformers.utils.hub as tf_hub
|
||||
|
||||
return tf_hub
|
||||
|
||||
|
||||
def test_cached_model_load_forces_offline_mode():
|
||||
original_hf = _hf_const().HF_HUB_OFFLINE
|
||||
original_tf = _tf_hub()._is_offline_mode
|
||||
|
||||
with model_load_progress("test-model", is_cached=True):
|
||||
assert _hf_const().HF_HUB_OFFLINE is True
|
||||
assert _tf_hub()._is_offline_mode is True
|
||||
|
||||
assert original_hf == _hf_const().HF_HUB_OFFLINE
|
||||
assert original_tf == _tf_hub()._is_offline_mode
|
||||
|
||||
|
||||
def test_uncached_model_load_does_not_force_offline_mode():
|
||||
original_hf = _hf_const().HF_HUB_OFFLINE
|
||||
|
||||
with model_load_progress("test-model", is_cached=False):
|
||||
assert original_hf == _hf_const().HF_HUB_OFFLINE
|
||||
|
||||
assert original_hf == _hf_const().HF_HUB_OFFLINE
|
||||
@@ -12,9 +12,11 @@ parallelism (e.g. ``pytest-xdist`` with ``--dist=loadfile``/``loadscope``);
|
||||
run this file serially.
|
||||
"""
|
||||
|
||||
import multiprocessing
|
||||
import os
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
@@ -24,6 +26,27 @@ sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
from utils.hf_offline_patch import force_offline_if_cached # noqa: E402
|
||||
|
||||
|
||||
def _nest_opposite_modes_in_subprocess(queue):
|
||||
"""Module-level so it's picklable for multiprocessing's spawn start method.
|
||||
|
||||
Runs in a fresh child process rather than a thread of the test process,
|
||||
so a real deadlock here (a regression of the guard this test exists for)
|
||||
can't leave the shared _offline_cv state corrupted for every other test
|
||||
in this run: the child either exits cleanly (guard raised) or gets
|
||||
terminated by the parent after the timeout, and either way the test
|
||||
process's own state was never touched.
|
||||
"""
|
||||
try:
|
||||
with force_offline_if_cached(False, "outer-uncached"), force_offline_if_cached(
|
||||
True, "inner-cached"
|
||||
):
|
||||
pass
|
||||
except Exception as exc:
|
||||
queue.put(("exc", type(exc).__name__, str(exc)))
|
||||
else:
|
||||
queue.put(("ok", None, None))
|
||||
|
||||
|
||||
def _hf_const():
|
||||
import huggingface_hub.constants as hf_const
|
||||
|
||||
@@ -114,5 +137,78 @@ def test_concurrent_threads_share_offline_window():
|
||||
assert original == _hf_const().HF_HUB_OFFLINE
|
||||
|
||||
|
||||
def test_uncached_load_never_observes_offline_flag_from_concurrent_cached_load():
|
||||
"""An uncached (network-needing) load must never inherit the forced
|
||||
offline mode of a concurrent, unrelated cached load — even when both
|
||||
start at nearly the same time.
|
||||
"""
|
||||
original = _hf_const().HF_HUB_OFFLINE
|
||||
observations: list[bool] = []
|
||||
errors: list[Exception] = []
|
||||
cached_entered = threading.Event()
|
||||
|
||||
def cached_load():
|
||||
try:
|
||||
with force_offline_if_cached(True, "cached"):
|
||||
cached_entered.set()
|
||||
time.sleep(0.2)
|
||||
except Exception as exc:
|
||||
errors.append(exc)
|
||||
|
||||
def uncached_load():
|
||||
try:
|
||||
assert cached_entered.wait(timeout=5), "cached thread never entered"
|
||||
with force_offline_if_cached(False, "uncached"):
|
||||
observations.append(_hf_const().HF_HUB_OFFLINE)
|
||||
except Exception as exc:
|
||||
errors.append(exc)
|
||||
|
||||
t_cached = threading.Thread(target=cached_load)
|
||||
t_uncached = threading.Thread(target=uncached_load)
|
||||
t_cached.start()
|
||||
t_uncached.start()
|
||||
t_cached.join(timeout=5)
|
||||
t_uncached.join(timeout=5)
|
||||
|
||||
assert not t_cached.is_alive(), "cached thread did not finish"
|
||||
assert not t_uncached.is_alive(), "uncached thread did not finish"
|
||||
assert not errors, errors
|
||||
assert observations == [False], "uncached load observed offline mode forced by a concurrent cached load"
|
||||
assert original == _hf_const().HF_HUB_OFFLINE
|
||||
|
||||
|
||||
def test_nesting_opposite_mode_on_same_thread_raises_instead_of_deadlocking():
|
||||
"""Nesting is_cached=True inside is_cached=False (or vice versa) on the
|
||||
same thread must raise immediately, not hang: the inner call's wait
|
||||
condition can only be cleared by the outer call's own exit, which can
|
||||
never run because it's blocked inside the inner call waiting for it.
|
||||
|
||||
Runs in a spawned child process with a bounded join, terminated if it's
|
||||
still alive after the timeout, so a regression fails this test instead of
|
||||
hanging the suite or leaving _offline_cv's shared state corrupted for
|
||||
every other test in this run.
|
||||
"""
|
||||
ctx = multiprocessing.get_context("spawn")
|
||||
queue = ctx.Queue()
|
||||
proc = ctx.Process(target=_nest_opposite_modes_in_subprocess, args=(queue,))
|
||||
proc.start()
|
||||
proc.join(timeout=5)
|
||||
|
||||
if proc.is_alive():
|
||||
proc.terminate()
|
||||
proc.join(timeout=2)
|
||||
if proc.is_alive():
|
||||
proc.kill()
|
||||
proc.join(timeout=2)
|
||||
pytest.fail(
|
||||
"nesting the opposite mode on the same thread hung instead of raising, "
|
||||
"this is the deadlock the per-thread mode-stack guard exists to prevent"
|
||||
)
|
||||
|
||||
kind, exc_type, exc_msg = queue.get(timeout=2)
|
||||
assert kind == "exc", (kind, exc_type, exc_msg)
|
||||
assert exc_type == "RuntimeError", (kind, exc_type, exc_msg)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v"])
|
||||
|
||||
Reference in New Issue
Block a user