mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 17:15:19 -07:00
fix(offline-guard): keep uncached loads from inheriting a concurrent cached load's offline window
This commit is contained in:
@@ -15,6 +15,7 @@ run this file serially.
|
|||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
import threading
|
import threading
|
||||||
|
import time
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -114,5 +115,45 @@ def test_concurrent_threads_share_offline_window():
|
|||||||
assert original == _hf_const().HF_HUB_OFFLINE
|
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
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
pytest.main([__file__, "-v"])
|
pytest.main([__file__, "-v"])
|
||||||
|
|||||||
@@ -22,9 +22,20 @@ logger = logging.getLogger(__name__)
|
|||||||
# ``transformers.utils.hub.is_offline_mode``) read the bools — not the env.
|
# ``transformers.utils.hub.is_offline_mode``) read the bools — not the env.
|
||||||
# We mutate the cached constants directly, guarded by a refcount so
|
# We mutate the cached constants directly, guarded by a refcount so
|
||||||
# concurrent inference threads share a single offline window safely.
|
# concurrent inference threads share a single offline window safely.
|
||||||
|
#
|
||||||
|
# That refcounted window is process-global, so an *uncached* load (which
|
||||||
|
# needs real network access) must never run while it's open — it would
|
||||||
|
# silently inherit HF_HUB_OFFLINE=True from an unrelated concurrent cached
|
||||||
|
# load and fail to fetch what it needs. Symmetrically, a new offline window
|
||||||
|
# must not open while an uncached load is in flight. Both directions wait on
|
||||||
|
# a condition variable rather than a plain lock: holding a lock for an
|
||||||
|
# entire model load (which can take minutes over the network) would also
|
||||||
|
# serialize unrelated *cached* loads against each other, which sharing the
|
||||||
|
# window is specifically meant to allow.
|
||||||
|
|
||||||
_offline_lock = threading.RLock()
|
_offline_cv = threading.Condition(threading.RLock())
|
||||||
_offline_refcount = 0
|
_offline_refcount = 0
|
||||||
|
_uncached_active = 0
|
||||||
_saved_env: Optional[str] = None
|
_saved_env: Optional[str] = None
|
||||||
_saved_hf_const: Optional[bool] = None
|
_saved_hf_const: Optional[bool] = None
|
||||||
_saved_transformers_const: Optional[bool] = None
|
_saved_transformers_const: Optional[bool] = None
|
||||||
@@ -40,19 +51,36 @@ def force_offline_if_cached(is_cached: bool, model_label: str = ""):
|
|||||||
so multiple concurrent inference threads share a single offline window
|
so multiple concurrent inference threads share a single offline window
|
||||||
and the last one to exit restores state.
|
and the last one to exit restores state.
|
||||||
|
|
||||||
If *is_cached* is ``False`` the block runs normally (network allowed).
|
If *is_cached* is ``False`` the block waits for any open offline window
|
||||||
|
to close first, then runs with network allowed — and blocks any new
|
||||||
|
offline window from opening until it's done, so it can never observe
|
||||||
|
(or be blamed for breaking) a concurrent cached load's forced-offline
|
||||||
|
state.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
is_cached: Whether the model weights are already on disk.
|
is_cached: Whether the model weights are already on disk.
|
||||||
model_label: Human-readable name used in log messages.
|
model_label: Human-readable name used in log messages.
|
||||||
"""
|
"""
|
||||||
|
global _offline_refcount, _uncached_active
|
||||||
|
global _saved_env, _saved_hf_const, _saved_transformers_const
|
||||||
|
|
||||||
if not is_cached:
|
if not is_cached:
|
||||||
yield
|
with _offline_cv:
|
||||||
|
while _offline_refcount > 0:
|
||||||
|
_offline_cv.wait()
|
||||||
|
_uncached_active += 1
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
with _offline_cv:
|
||||||
|
_uncached_active -= 1
|
||||||
|
_offline_cv.notify_all()
|
||||||
return
|
return
|
||||||
|
|
||||||
global _offline_refcount, _saved_env, _saved_hf_const, _saved_transformers_const
|
with _offline_cv:
|
||||||
|
while _uncached_active > 0:
|
||||||
|
_offline_cv.wait()
|
||||||
|
|
||||||
with _offline_lock:
|
|
||||||
if _offline_refcount == 0:
|
if _offline_refcount == 0:
|
||||||
# Snapshot prior state, apply new state, roll back on *any*
|
# Snapshot prior state, apply new state, roll back on *any*
|
||||||
# failure. Catching only ImportError here would let a partially
|
# failure. Catching only ImportError here would let a partially
|
||||||
@@ -116,7 +144,7 @@ def force_offline_if_cached(is_cached: bool, model_label: str = ""):
|
|||||||
try:
|
try:
|
||||||
yield
|
yield
|
||||||
finally:
|
finally:
|
||||||
with _offline_lock:
|
with _offline_cv:
|
||||||
_offline_refcount -= 1
|
_offline_refcount -= 1
|
||||||
if _offline_refcount == 0:
|
if _offline_refcount == 0:
|
||||||
if _saved_env is not None:
|
if _saved_env is not None:
|
||||||
@@ -140,6 +168,7 @@ def force_offline_if_cached(is_cached: bool, model_label: str = ""):
|
|||||||
_saved_env = None
|
_saved_env = None
|
||||||
_saved_hf_const = None
|
_saved_hf_const = None
|
||||||
_saved_transformers_const = None
|
_saved_transformers_const = None
|
||||||
|
_offline_cv.notify_all()
|
||||||
|
|
||||||
|
|
||||||
_mistral_regex_patched = False
|
_mistral_regex_patched = False
|
||||||
|
|||||||
Reference in New Issue
Block a user