mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 17:15:19 -07:00
fix(offline-guard): make every cache check cover the files the forced-offline load reads
The load body now runs with HF_HUB_OFFLINE forced when _is_model_cached() reports True, so a snapshot that holds the weights but not the small files the loader also opens would fail hard instead of fetching them. Verified against chatterbox-tts 0.1.7 (mtl_tts.py allow_patterns, tts_turbo.py from_local) and the TADA loader's unsloth/Llama-3.2-1B tokenizer download; list those files in the required_files checks. Also note in the changelog why the 0.4.5 removal of this guard no longer applies.
This commit is contained in:
+10
-1
@@ -13,7 +13,16 @@
|
||||
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.
|
||||
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, every backend's cache check
|
||||
lists the full set of files its load reads (Chatterbox tokenizer/conds, TADA's Llama
|
||||
tokenizer mirror), so a partially downloaded snapshot reports "not cached" and downloads
|
||||
online instead.
|
||||
|
||||
### Linux
|
||||
|
||||
|
||||
@@ -29,11 +29,16 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
CHATTERBOX_HF_REPO = "ResembleAI/chatterbox"
|
||||
|
||||
# Files that must be present for the multilingual model
|
||||
# Every file ChatterboxMultilingualTTS.from_pretrained() downloads (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.
|
||||
_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,18 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
CHATTERBOX_TURBO_HF_REPO = "ResembleAI/chatterbox-turbo"
|
||||
|
||||
# Files that must be present for the turbo model
|
||||
# Every file ChatterboxTurboTTS.from_local() reads (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.
|
||||
_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*"],
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user