Scope the video base repo download to the files the pipeline loads

A bare from_pretrained snapshot of Lightricks/LTX-2 pulls the whole 314 GB
repo: 170 GB of packaged root checkpoints and a second 50 GB text-encoder
shard set, when the pipeline reads about 93 GB. Build the needed file list
once (shared with the progress estimate so the two cannot disagree),
download it per file with cancellation, and hand from_pretrained the local
snapshot dir. Clamp the progress counter to the estimate so stale cache
blobs can no longer report over 100 percent.
This commit is contained in:
Daniel Han 2026-07-05 03:17:55 +00:00
commit d442e00407
2 changed files with 150 additions and 16 deletions

View file

@ -308,8 +308,8 @@ class VideoBackend:
self._loading.expected_bytes = expected
# The GGUF/single-file checkpoint downloads outside the lock so an
# unload/eviction can preempt the multi-GB pull; the pipeline
# companions download inside from_pretrained (which resumes from the
# cache, so a cancelled pull costs nothing).
# companions pre-download the same way (scoped file list, cancellable,
# resumes from the cache so a cancelled pull costs nothing).
if kwargs.get("gguf_filename") and not Path(kwargs["repo_id"]).expanduser().exists():
from utils.hf_xet_fallback import hf_hub_download_with_xet_fallback
hf_hub_download_with_xet_fallback(
@ -318,6 +318,9 @@ class VideoBackend:
kwargs.get("hf_token"),
cancel_event = self._cancel_event,
)
kwargs["_base_local_dir"] = self._predownload_base(
base, kwargs.get("hf_token"), kind
)
self.load_pipeline(**kwargs)
with self._lock:
if self._load_token == token:
@ -332,6 +335,32 @@ class VideoBackend:
if self._load_token == token and self._loading is not None:
self._loading.error = redact_native_paths(str(exc))
@staticmethod
def _base_download_files(info: Any, kind: str) -> list[tuple[str, int]]:
"""The (rfilename, size) list a load actually needs from the base repo.
Single source of truth for the progress estimate AND the scoped pre-download,
so the two can never disagree. Excluded on purpose:
- root-level packaged checkpoints (ComfyUI-style singles; 170 GB of the LTX-2
repo) -- the diffusers pipeline only reads per-component subfolders;
- the duplicate ``text_encoder/diffusion_pytorch_model*`` shard set (the LTX-2
base repo ships its text encoder twice; transformers loads the ``model-*``
naming via the shard index);
- ``transformer/`` when a GGUF/single-file checkpoint replaces the DiT."""
files: list[tuple[str, int]] = []
for sibling in info.siblings or []:
name, size = sibling.rfilename, sibling.size or 0
if not name.endswith((".safetensors", ".json", ".model", ".txt")):
continue
if "/" not in name and name.endswith(".safetensors"):
continue
if kind != "pipeline" and name.startswith("transformer/"):
continue
if name.startswith("text_encoder/diffusion_pytorch_model"):
continue
files.append((name, int(size)))
return files
def _estimate_download_bytes(
self,
repo_id: str,
@ -353,22 +382,48 @@ class VideoBackend:
total += int(sibling.size)
if base and not Path(base).expanduser().exists():
info = api.model_info(base, files_metadata = True)
for sibling in info.siblings or []:
name, size = sibling.rfilename, sibling.size or 0
if not name.endswith((".safetensors", ".json", ".model", ".txt")):
continue
# A GGUF/single-file load replaces the base repo's DiT, and the
# LTX-2 base repo ships its text encoder TWICE (two shard
# namings); count only what from_pretrained will pull.
if kind != "pipeline" and name.startswith("transformer/"):
continue
if name.startswith("text_encoder/diffusion_pytorch_model"):
continue
total += int(size)
total += sum(size for _, size in self._base_download_files(info, kind))
return total or None
except Exception: # noqa: BLE001 -- progress totals are best-effort only
return None
def _predownload_base(self, base: str, hf_token: Optional[str], kind: str) -> Optional[str]:
"""Pull exactly the base-repo files the load needs; return the local snapshot dir.
A bare ``from_pretrained(repo_id)`` snapshot of Lightricks/LTX-2 downloads the
whole 314 GB repo (root packaged checkpoints plus a second 50 GB text-encoder
shard set) when ~93 GB is used. Downloading the scoped file list ourselves is
also cancellable per file, and handing the local dir to from_pretrained skips
diffusers' own expected-files sweep. None -> caller keeps the hub id (local
path, non-diffusers layout, or any metadata failure: from_pretrained then
resolves the repo exactly as before)."""
try:
if not base or Path(base).expanduser().exists():
return None
from huggingface_hub import HfApi
info = HfApi(token = hf_token or None).model_info(base, files_metadata = True)
files = self._base_download_files(info, kind)
if not any(name == "model_index.json" for name, _ in files):
return None
from utils.hf_xet_fallback import hf_hub_download_with_xet_fallback
snapshot_root: Optional[Path] = None
for name, _ in files:
local = Path(
hf_hub_download_with_xet_fallback(
base, name, hf_token, cancel_event = self._cancel_event
)
)
if name == "model_index.json":
snapshot_root = local.parent
return str(snapshot_root) if snapshot_root is not None else None
except Exception as exc: # noqa: BLE001 -- fall back to from_pretrained's own pull
if self._cancel_event.is_set():
raise
logger.warning("video.predownload_fallback: %s", exc)
return None
def _cache_bytes(self, repo_id: Optional[str]) -> int:
"""Bytes of ``repo_id`` currently in the HF blob cache (progress polling)."""
if not repo_id:
@ -397,6 +452,11 @@ class VideoBackend:
phase = "downloading"
if expected and downloaded >= expected:
phase = "finalizing"
# The cache scan counts every blob of the repo(s), including files a
# previous (or broader) pull left behind that this load never reads, so
# the raw counter can exceed the scoped estimate. Clamp: everything the
# load needs is present, which is what the bar reports.
downloaded = expected
return _progress(
phase,
downloaded_bytes = int(downloaded),
@ -420,6 +480,7 @@ class VideoBackend:
transformer_cache_threshold: Optional[float] = None,
model_kind: Optional[str] = None,
_load_token: Optional[int] = None,
_base_local_dir: Optional[str] = None,
) -> dict[str, Any]:
import diffusers
import torch
@ -506,7 +567,10 @@ class VideoBackend:
if hf_token:
pipe_kwargs["token"] = hf_token
if kind == "pipeline":
pipe = pipeline_cls.from_pretrained(repo_id, **pipe_kwargs)
# The pre-downloaded snapshot dir keeps from_pretrained off the hub (its
# own snapshot sweep would also pull the repo's packaged root checkpoints
# and duplicate text-encoder shards); hub id when pre-download was skipped.
pipe = pipeline_cls.from_pretrained(_base_local_dir or repo_id, **pipe_kwargs)
else:
transformer_cls = getattr(diffusers, fam.transformer_class)
# checkpoint_path was already resolved (and downloaded) by the memory
@ -536,7 +600,9 @@ class VideoBackend:
)
else:
transformer = transformer_cls.from_single_file(str(checkpoint_path), **sf_kwargs)
pipe = pipeline_cls.from_pretrained(base, transformer = transformer, **pipe_kwargs)
pipe = pipeline_cls.from_pretrained(
_base_local_dir or base, transformer = transformer, **pipe_kwargs
)
if _load_token is not None and _load_token != self._load_token:
del pipe

View file

@ -441,3 +441,71 @@ def test_generate_progress_and_cancel_idle(fake_runtime):
def test_singleton():
assert get_video_backend() is get_video_backend()
def _sibling(name, size):
return types.SimpleNamespace(rfilename = name, size = size)
_LTX2_SIBLINGS = [
_sibling("model_index.json", 10),
_sibling("ltx-2-19b-packaged-fp8.safetensors", 170),
_sibling("transformer/config.json", 1),
_sibling("transformer/diffusion_pytorch_model-00001-of-00002.safetensors", 20),
_sibling("transformer/diffusion_pytorch_model-00002-of-00002.safetensors", 18),
_sibling("text_encoder/model-00001-of-00002.safetensors", 25),
_sibling("text_encoder/model-00002-of-00002.safetensors", 25),
_sibling("text_encoder/diffusion_pytorch_model-00001-of-00002.safetensors", 25),
_sibling("text_encoder/diffusion_pytorch_model-00002-of-00002.safetensors", 25),
_sibling("vae/diffusion_pytorch_model.safetensors", 3),
_sibling("tokenizer/tokenizer.model", 1),
_sibling("assets/example.mp4", 500),
]
def test_base_download_files_scopes_pipeline_pull():
# A pipeline load skips the packaged root checkpoint, the duplicate
# text-encoder shard naming, and non-weight assets -- and keeps everything else.
info = types.SimpleNamespace(siblings = _LTX2_SIBLINGS)
files = dict(VideoBackend._base_download_files(info, "pipeline"))
assert "ltx-2-19b-packaged-fp8.safetensors" not in files
assert "text_encoder/diffusion_pytorch_model-00001-of-00002.safetensors" not in files
assert "assets/example.mp4" not in files
assert files["text_encoder/model-00001-of-00002.safetensors"] == 25
assert files["transformer/diffusion_pytorch_model-00001-of-00002.safetensors"] == 20
assert sum(files.values()) == 10 + 1 + 20 + 18 + 25 + 25 + 3 + 1
def test_base_download_files_gguf_drops_transformer():
# A GGUF/single-file checkpoint replaces the DiT: the base transformer never pulls.
info = types.SimpleNamespace(siblings = _LTX2_SIBLINGS)
names = [n for n, _ in VideoBackend._base_download_files(info, "gguf")]
assert not any(n.startswith("transformer/") for n in names)
assert "text_encoder/model-00001-of-00002.safetensors" in names
def test_load_progress_clamps_overshoot(fake_runtime, monkeypatch):
# The cache scan counts blobs a broader previous pull left behind; the reported
# counter must never exceed the scoped estimate (no "282 GB of 263 GB").
backend = VideoBackend()
backend._loading = types.SimpleNamespace(
repo_id = "Lightricks/LTX-2", base_repo = None, expected_bytes = 100, error = None
)
monkeypatch.setattr(VideoBackend, "_cache_bytes", lambda self, repo: 150)
progress = backend.load_progress()
assert progress["phase"] == "finalizing"
assert progress["downloaded_bytes"] == 100
assert progress["expected_bytes"] == 100
def test_pipeline_load_uses_predownloaded_dir(fake_runtime, tmp_path):
# When the scoped pre-download produced a local snapshot, from_pretrained must
# receive that dir (keeping diffusers' own broader snapshot sweep off the hub).
backend = VideoBackend()
backend.load_pipeline(
"Lightricks/LTX-2",
model_kind = "pipeline",
_base_local_dir = str(tmp_path),
)
assert _FakePipeline.last["base"] == str(tmp_path)
backend.unload()