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:
capy-ai-staging[bot]
2026-10-04 00:05:46 +00:00
committed by GitHub
8 changed files with 351 additions and 90 deletions
+20
View File
@@ -32,6 +32,26 @@
- Voice generation requests that omit `engine` now honor the selected profile's
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
- **ROCm setup works on Linux AMD systems.** Docker ROCm builds now keep PyTorch
+3 -1
View File
@@ -15,6 +15,7 @@ from typing import Callable, List, Optional, Tuple
import numpy as np
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.hf_progress import HFProgressTracker, create_hf_progress_callback
from ..utils.tasks import get_task_manager
@@ -319,7 +320,8 @@ def model_load_progress(
)
try:
yield tracker_context
with force_offline_if_cached(is_cached, model_name):
yield tracker_context
except Exception as e:
# Report error to both managers
progress_manager.mark_error(model_name, str(e))
+7 -1
View File
@@ -29,11 +29,17 @@ logger = logging.getLogger(__name__)
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 = [
"t3_mtl23ls_v2.safetensors",
"s3gen.pt",
"ve.pt",
"grapheme_mtl_merged_expanded_v1.json",
"conds.pt",
"Cangjie5_TC.json",
]
+9 -1
View File
@@ -29,11 +29,19 @@ logger = logging.getLogger(__name__)
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 = [
"t3_turbo_v1.safetensors",
"s3gen_meanflow.safetensors",
"ve.safetensors",
"tokenizer_config.json",
"vocab.json",
"merges.txt",
"conds.pt",
]
+12 -2
View File
@@ -35,6 +35,9 @@ logger = logging.getLogger(__name__)
# HuggingFace repos
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_3B_ML_REPO = "HumeAI/tada-3b-ml"
@@ -52,6 +55,12 @@ _TADA_CODEC_WEIGHT_FILES = [
"encoder/model.safetensors",
]
_TADA_TOKENIZER_FILES = [
"tokenizer.json",
"tokenizer_config.json",
"special_tokens_map.json",
]
class HumeTadaBackend:
"""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)
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)
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:
"""Load the TADA model and encoder."""
@@ -140,7 +150,7 @@ class HumeTadaBackend:
# local cache path so we can point TADA at it directly.
logger.info("Downloading Llama tokenizer (ungated mirror)...")
tokenizer_path = snapshot_download(
repo_id="unsloth/Llama-3.2-1B",
repo_id=TADA_TOKENIZER_REPO,
token=None,
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
+96
View File
@@ -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"])
+149 -85
View File
@@ -22,13 +22,42 @@ logger = logging.getLogger(__name__)
# ``transformers.utils.hub.is_offline_mode``) read the bools — not the env.
# We mutate the cached constants directly, guarded by a refcount so
# 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
_uncached_active = 0
_saved_env: Optional[str] = None
_saved_hf_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
def force_offline_if_cached(is_cached: bool, model_label: str = ""):
@@ -40,106 +69,141 @@ def force_offline_if_cached(is_cached: bool, model_label: str = ""):
so multiple concurrent inference threads share a single offline window
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:
is_cached: Whether the model weights are already on disk.
model_label: Human-readable name used in log messages.
"""
if not is_cached:
yield
return
global _offline_refcount, _saved_env, _saved_hf_const, _saved_transformers_const
with _offline_lock:
if _offline_refcount == 0:
# Snapshot prior state, apply new state, roll back on *any*
# failure. Catching only ImportError here would let a partially
# broken install (RuntimeError, AttributeError from a half-init
# module, etc.) leave the cached HF constants mutated without
# bumping the refcount — a persistent offline leak that outlives
# the process and is miserable to debug.
prev_env = os.environ.get("HF_HUB_OFFLINE")
prev_hf: Optional[bool] = None
prev_tf: Optional[bool] = None
try:
try:
import huggingface_hub.constants as hf_const
prev_hf = hf_const.HF_HUB_OFFLINE
hf_const.HF_HUB_OFFLINE = True
except ImportError:
prev_hf = None
try:
import transformers.utils.hub as tf_hub
prev_tf = tf_hub._is_offline_mode
tf_hub._is_offline_mode = True
except ImportError:
prev_tf = None
os.environ["HF_HUB_OFFLINE"] = "1"
except BaseException:
# Roll back whatever we already changed, then re-raise so
# the caller sees the real failure.
if prev_hf is not None:
try:
import huggingface_hub.constants as hf_const
hf_const.HF_HUB_OFFLINE = prev_hf
except ImportError:
pass
if prev_tf is not None:
try:
import transformers.utils.hub as tf_hub
tf_hub._is_offline_mode = prev_tf
except ImportError:
pass
if prev_env is not None:
os.environ["HF_HUB_OFFLINE"] = prev_env
else:
os.environ.pop("HF_HUB_OFFLINE", None)
raise
_saved_env = prev_env
_saved_hf_const = prev_hf
_saved_transformers_const = prev_tf
logger.info(
"[offline-guard] %s is cached — forcing offline mode",
model_label or "model",
)
_offline_refcount += 1
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:
yield
finally:
with _offline_lock:
_offline_refcount -= 1
if not is_cached:
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
with _offline_cv:
while _uncached_active > 0:
_offline_cv.wait()
if _offline_refcount == 0:
if _saved_env is not None:
os.environ["HF_HUB_OFFLINE"] = _saved_env
else:
os.environ.pop("HF_HUB_OFFLINE", None)
if _saved_hf_const is not None:
# Snapshot prior state, apply new state, roll back on *any*
# failure. Catching only ImportError here would let a partially
# broken install (RuntimeError, AttributeError from a half-init
# module, etc.) leave the cached HF constants mutated without
# bumping the refcount — a persistent offline leak that outlives
# the process and is miserable to debug.
prev_env = os.environ.get("HF_HUB_OFFLINE")
prev_hf: Optional[bool] = None
prev_tf: Optional[bool] = None
try:
try:
import huggingface_hub.constants as hf_const
hf_const.HF_HUB_OFFLINE = _saved_hf_const
prev_hf = hf_const.HF_HUB_OFFLINE
hf_const.HF_HUB_OFFLINE = True
except ImportError:
pass
if _saved_transformers_const is not None:
prev_hf = None
try:
import transformers.utils.hub as tf_hub
tf_hub._is_offline_mode = _saved_transformers_const
prev_tf = tf_hub._is_offline_mode
tf_hub._is_offline_mode = True
except ImportError:
pass
_saved_env = None
_saved_hf_const = None
_saved_transformers_const = None
prev_tf = None
os.environ["HF_HUB_OFFLINE"] = "1"
except BaseException:
# Roll back whatever we already changed, then re-raise so
# the caller sees the real failure.
if prev_hf is not None:
try:
import huggingface_hub.constants as hf_const
hf_const.HF_HUB_OFFLINE = prev_hf
except ImportError:
pass
if prev_tf is not None:
try:
import transformers.utils.hub as tf_hub
tf_hub._is_offline_mode = prev_tf
except ImportError:
pass
if prev_env is not None:
os.environ["HF_HUB_OFFLINE"] = prev_env
else:
os.environ.pop("HF_HUB_OFFLINE", None)
raise
_saved_env = prev_env
_saved_hf_const = prev_hf
_saved_transformers_const = prev_tf
logger.info(
"[offline-guard] %s is cached — forcing offline mode",
model_label or "model",
)
_offline_refcount += 1
try:
yield
finally:
with _offline_cv:
_offline_refcount -= 1
if _offline_refcount == 0:
if _saved_env is not None:
os.environ["HF_HUB_OFFLINE"] = _saved_env
else:
os.environ.pop("HF_HUB_OFFLINE", None)
if _saved_hf_const is not None:
try:
import huggingface_hub.constants as hf_const
hf_const.HF_HUB_OFFLINE = _saved_hf_const
except ImportError:
pass
if _saved_transformers_const is not None:
try:
import transformers.utils.hub as tf_hub
tf_hub._is_offline_mode = _saved_transformers_const
except ImportError:
pass
_saved_env = None
_saved_hf_const = None
_saved_transformers_const = None
_offline_cv.notify_all()
finally:
stack.pop()
_mistral_regex_patched = False