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:
@@ -32,6 +32,26 @@
|
|||||||
- Voice generation requests that omit `engine` now honor the selected profile's
|
- Voice generation requests that omit `engine` now honor the selected profile's
|
||||||
configured engine instead of silently defaulting to Qwen.
|
configured engine instead of silently defaulting to Qwen.
|
||||||
|
|
||||||
|
### Reliability
|
||||||
|
|
||||||
|
- **Cached models no longer retry HuggingFace when offline.** Loading a fully-downloaded
|
||||||
|
model now forces offline mode for the duration of the load, so it skips the network HEAD
|
||||||
|
request (and its 5-retry backoff) for every config file — `config.json`,
|
||||||
|
`generation_config.json`, and the rest — instead of retrying each one in sequence before the
|
||||||
|
app becomes ready. This reinstates the load-time `force_offline_if_cached` guard that 0.4.5
|
||||||
|
([#530](https://github.com/jamiepine/voicebox/pull/530)) removed: that removal was a hotfix
|
||||||
|
for the `_patch_mistral_regex` crash ([#526](https://github.com/jamiepine/voicebox/issues/526)),
|
||||||
|
which the wrapper installed in the same release now catches at the source, so the guard no
|
||||||
|
longer trips it. The per-file HEAD retries from
|
||||||
|
[#434](https://github.com/jamiepine/voicebox/issues/434) were never covered by that wrapper.
|
||||||
|
Because a load now fails hard offline when any file is missing, the Chatterbox, Chatterbox
|
||||||
|
Turbo, and TADA cache checks were extended to the small files their loaders also read
|
||||||
|
(tokenizer files, `conds.pt`, TADA's Llama tokenizer mirror) so a snapshot missing one of
|
||||||
|
them reports "not cached" and downloads online instead. The other engines still gate on
|
||||||
|
their weight files (plus `config.json` for Kokoro); transformers-based loaders fetch config
|
||||||
|
before weights, so a weights-present cache normally holds the rest, but that is an
|
||||||
|
assumption, not a check.
|
||||||
|
|
||||||
### Linux
|
### Linux
|
||||||
|
|
||||||
- **ROCm setup works on Linux AMD systems.** Docker ROCm builds now keep PyTorch
|
- **ROCm setup works on Linux AMD systems.** Docker ROCm builds now keep PyTorch
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ from typing import Callable, List, Optional, Tuple
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
from ..utils.audio import normalize_audio, load_audio
|
from ..utils.audio import normalize_audio, load_audio
|
||||||
|
from ..utils.hf_offline_patch import force_offline_if_cached
|
||||||
from ..utils.progress import get_progress_manager
|
from ..utils.progress import get_progress_manager
|
||||||
from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback
|
from ..utils.hf_progress import HFProgressTracker, create_hf_progress_callback
|
||||||
from ..utils.tasks import get_task_manager
|
from ..utils.tasks import get_task_manager
|
||||||
@@ -319,6 +320,7 @@ def model_load_progress(
|
|||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
with force_offline_if_cached(is_cached, model_name):
|
||||||
yield tracker_context
|
yield tracker_context
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
# Report error to both managers
|
# Report error to both managers
|
||||||
|
|||||||
@@ -29,11 +29,17 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
CHATTERBOX_HF_REPO = "ResembleAI/chatterbox"
|
CHATTERBOX_HF_REPO = "ResembleAI/chatterbox"
|
||||||
|
|
||||||
# Files that must be present for the multilingual model
|
# The files ChatterboxMultilingualTTS.from_pretrained() downloads, as of
|
||||||
|
# chatterbox-tts 0.1.7 (mtl_tts.py allow_patterns). The load runs with HF
|
||||||
|
# offline mode forced when this reports cached, so a partial snapshot must not
|
||||||
|
# count as cached -- if upstream adds a file to that list, add it here too.
|
||||||
_MTL_WEIGHT_FILES = [
|
_MTL_WEIGHT_FILES = [
|
||||||
"t3_mtl23ls_v2.safetensors",
|
"t3_mtl23ls_v2.safetensors",
|
||||||
"s3gen.pt",
|
"s3gen.pt",
|
||||||
"ve.pt",
|
"ve.pt",
|
||||||
|
"grapheme_mtl_merged_expanded_v1.json",
|
||||||
|
"conds.pt",
|
||||||
|
"Cangjie5_TC.json",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -29,11 +29,19 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
CHATTERBOX_TURBO_HF_REPO = "ResembleAI/chatterbox-turbo"
|
CHATTERBOX_TURBO_HF_REPO = "ResembleAI/chatterbox-turbo"
|
||||||
|
|
||||||
# Files that must be present for the turbo model
|
# The files ChatterboxTurboTTS.from_local() reads, as of chatterbox-tts 0.1.7:
|
||||||
|
# the three weight files, the GPT-2 tokenizer files AutoTokenizer needs, and
|
||||||
|
# the built-in voice. The load runs with HF offline mode forced when this
|
||||||
|
# reports cached, so a partial snapshot must not count as cached -- if upstream
|
||||||
|
# starts reading another file, add it here too.
|
||||||
_TURBO_WEIGHT_FILES = [
|
_TURBO_WEIGHT_FILES = [
|
||||||
"t3_turbo_v1.safetensors",
|
"t3_turbo_v1.safetensors",
|
||||||
"s3gen_meanflow.safetensors",
|
"s3gen_meanflow.safetensors",
|
||||||
"ve.safetensors",
|
"ve.safetensors",
|
||||||
|
"tokenizer_config.json",
|
||||||
|
"vocab.json",
|
||||||
|
"merges.txt",
|
||||||
|
"conds.pt",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -35,6 +35,9 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
# HuggingFace repos
|
# HuggingFace repos
|
||||||
TADA_CODEC_REPO = "HumeAI/tada-codec"
|
TADA_CODEC_REPO = "HumeAI/tada-codec"
|
||||||
|
# TADA hardcodes the gated meta-llama/Llama-3.2-1B tokenizer; we load it from
|
||||||
|
# this ungated mirror instead (see load_model).
|
||||||
|
TADA_TOKENIZER_REPO = "unsloth/Llama-3.2-1B"
|
||||||
TADA_1B_REPO = "HumeAI/tada-1b"
|
TADA_1B_REPO = "HumeAI/tada-1b"
|
||||||
TADA_3B_ML_REPO = "HumeAI/tada-3b-ml"
|
TADA_3B_ML_REPO = "HumeAI/tada-3b-ml"
|
||||||
|
|
||||||
@@ -52,6 +55,12 @@ _TADA_CODEC_WEIGHT_FILES = [
|
|||||||
"encoder/model.safetensors",
|
"encoder/model.safetensors",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
_TADA_TOKENIZER_FILES = [
|
||||||
|
"tokenizer.json",
|
||||||
|
"tokenizer_config.json",
|
||||||
|
"special_tokens_map.json",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
class HumeTadaBackend:
|
class HumeTadaBackend:
|
||||||
"""HumeAI TADA TTS backend for high-quality voice cloning."""
|
"""HumeAI TADA TTS backend for high-quality voice cloning."""
|
||||||
@@ -80,7 +89,8 @@ class HumeTadaBackend:
|
|||||||
repo = TADA_MODEL_REPOS.get(model_size, TADA_1B_REPO)
|
repo = TADA_MODEL_REPOS.get(model_size, TADA_1B_REPO)
|
||||||
model_cached = is_model_cached(repo, required_files=_TADA_MODEL_WEIGHT_FILES)
|
model_cached = is_model_cached(repo, required_files=_TADA_MODEL_WEIGHT_FILES)
|
||||||
codec_cached = is_model_cached(TADA_CODEC_REPO, required_files=_TADA_CODEC_WEIGHT_FILES)
|
codec_cached = is_model_cached(TADA_CODEC_REPO, required_files=_TADA_CODEC_WEIGHT_FILES)
|
||||||
return model_cached and codec_cached
|
tokenizer_cached = is_model_cached(TADA_TOKENIZER_REPO, required_files=_TADA_TOKENIZER_FILES)
|
||||||
|
return model_cached and codec_cached and tokenizer_cached
|
||||||
|
|
||||||
async def load_model(self, model_size: str = "1B") -> None:
|
async def load_model(self, model_size: str = "1B") -> None:
|
||||||
"""Load the TADA model and encoder."""
|
"""Load the TADA model and encoder."""
|
||||||
@@ -140,7 +150,7 @@ class HumeTadaBackend:
|
|||||||
# local cache path so we can point TADA at it directly.
|
# local cache path so we can point TADA at it directly.
|
||||||
logger.info("Downloading Llama tokenizer (ungated mirror)...")
|
logger.info("Downloading Llama tokenizer (ungated mirror)...")
|
||||||
tokenizer_path = snapshot_download(
|
tokenizer_path = snapshot_download(
|
||||||
repo_id="unsloth/Llama-3.2-1B",
|
repo_id=TADA_TOKENIZER_REPO,
|
||||||
token=None,
|
token=None,
|
||||||
allow_patterns=["tokenizer*", "special_tokens*"],
|
allow_patterns=["tokenizer*", "special_tokens*"],
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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.
|
run this file serially.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import multiprocessing
|
||||||
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
|
||||||
@@ -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
|
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():
|
def _hf_const():
|
||||||
import huggingface_hub.constants as 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
|
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__":
|
if __name__ == "__main__":
|
||||||
pytest.main([__file__, "-v"])
|
pytest.main([__file__, "-v"])
|
||||||
|
|||||||
@@ -22,13 +22,42 @@ 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
|
||||||
|
|
||||||
|
# Per-thread stack of the modes (`is_cached` values) a thread currently holds
|
||||||
|
# open. Same-mode nesting on one thread is fine (test_nested_contexts_respect_
|
||||||
|
# refcount relies on it) and re-enters `_offline_cv`'s RLock without blocking.
|
||||||
|
# Opposite-mode nesting on the *same* thread would deadlock instead: the inner
|
||||||
|
# call's `while ...: _offline_cv.wait()` would wait on a condition only some
|
||||||
|
# *other* thread's exit can clear, but `Condition.wait()` on an RLock releases
|
||||||
|
# only one recursion level, so the outer call's level stays held and no other
|
||||||
|
# thread can ever acquire it to notify. Detect and refuse instead of hanging.
|
||||||
|
_thread_modes = threading.local()
|
||||||
|
|
||||||
|
|
||||||
|
def _mode_stack() -> list:
|
||||||
|
try:
|
||||||
|
return _thread_modes.stack
|
||||||
|
except AttributeError:
|
||||||
|
_thread_modes.stack = []
|
||||||
|
return _thread_modes.stack
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def force_offline_if_cached(is_cached: bool, model_label: str = ""):
|
def force_offline_if_cached(is_cached: bool, model_label: str = ""):
|
||||||
@@ -40,19 +69,51 @@ 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.
|
||||||
|
|
||||||
|
Nesting calls with the *same* ``is_cached`` value on one thread is
|
||||||
|
supported. Nesting the opposite value on the same thread raises
|
||||||
|
``RuntimeError`` instead of deadlocking (see ``_thread_modes`` above).
|
||||||
|
|
||||||
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
|
||||||
|
|
||||||
|
stack = _mode_stack()
|
||||||
|
if stack and stack[-1] != is_cached:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"force_offline_if_cached({is_cached!r}, {model_label!r}) was called "
|
||||||
|
f"while this thread already holds a force_offline_if_cached({stack[-1]!r}, ...) "
|
||||||
|
"context open. Nesting the opposite mode on the same thread would "
|
||||||
|
"deadlock rather than raise; nest same-mode calls only, or run the "
|
||||||
|
"other mode on a different thread."
|
||||||
|
)
|
||||||
|
stack.append(is_cached)
|
||||||
|
try:
|
||||||
if not is_cached:
|
if not is_cached:
|
||||||
|
with _offline_cv:
|
||||||
|
while _offline_refcount > 0:
|
||||||
|
_offline_cv.wait()
|
||||||
|
_uncached_active += 1
|
||||||
|
try:
|
||||||
yield
|
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 +177,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 +201,9 @@ 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()
|
||||||
|
finally:
|
||||||
|
stack.pop()
|
||||||
|
|
||||||
|
|
||||||
_mistral_regex_patched = False
|
_mistral_regex_patched = False
|
||||||
|
|||||||
Reference in New Issue
Block a user