mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 01:25:18 -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
|
||||
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
|
||||
|
||||
@@ -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,6 +320,7 @@ def model_load_progress(
|
||||
)
|
||||
|
||||
try:
|
||||
with force_offline_if_cached(is_cached, model_name):
|
||||
yield tracker_context
|
||||
except Exception as e:
|
||||
# Report error to both managers
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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"])
|
||||
|
||||
@@ -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,19 +69,51 @@ 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.
|
||||
"""
|
||||
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:
|
||||
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
|
||||
|
||||
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:
|
||||
# Snapshot prior state, apply new state, roll back on *any*
|
||||
# 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:
|
||||
yield
|
||||
finally:
|
||||
with _offline_lock:
|
||||
with _offline_cv:
|
||||
_offline_refcount -= 1
|
||||
if _offline_refcount == 0:
|
||||
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_hf_const = None
|
||||
_saved_transformers_const = None
|
||||
_offline_cv.notify_all()
|
||||
finally:
|
||||
stack.pop()
|
||||
|
||||
|
||||
_mistral_regex_patched = False
|
||||
|
||||
Reference in New Issue
Block a user