From e2fd90783e087947c90b463a174e2da4c0ffe486 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sun, 5 Jul 2026 04:38:28 +0000 Subject: [PATCH] Shrink the LTX-2.3 base pull to the components the assembly reads A 2.3 checkpoint (GGUF or single file) carries the DiT and, with its extras files, the connectors, both VAEs and the vocoder; only the 2.0 base repo's scheduler, text encoder and tokenizer are read. Detect 2.3 from the checkpoint header after the pull, re-estimate, and scope the base pre-download accordingly (about 6 GB less per fresh install). --- studio/backend/core/inference/video.py | 73 ++++++++++++++++++---- studio/backend/tests/test_video_backend.py | 20 ++++++ 2 files changed, 82 insertions(+), 11 deletions(-) diff --git a/studio/backend/core/inference/video.py b/studio/backend/core/inference/video.py index f00ccdee6b..d91cc62737 100644 --- a/studio/backend/core/inference/video.py +++ b/studio/backend/core/inference/video.py @@ -310,15 +310,49 @@ class VideoBackend: # unload/eviction can preempt the multi-GB pull; the pipeline # companions pre-download the same way (scoped file list, cancellable, # resumes from the cache so a cancelled pull costs nothing). + checkpoint_local: Optional[Path] = None 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( - kwargs["repo_id"], - kwargs["gguf_filename"], - kwargs.get("hf_token"), - cancel_event = self._cancel_event, + checkpoint_local = Path( + hf_hub_download_with_xet_fallback( + kwargs["repo_id"], + kwargs["gguf_filename"], + kwargs.get("hf_token"), + cancel_event = self._cancel_event, + ) ) - kwargs["_base_local_dir"] = self._predownload_base(base, kwargs.get("hf_token"), kind) + # An LTX-2.3 checkpoint replaces the base VAEs/vocoder/connectors too, so + # its base pull shrinks to scheduler + text encoder + tokenizer; the + # estimate is recomputed to match (detectable only once the checkpoint + # header is on disk, hence after the pull above). + ltx23 = False + if fam is not None and fam.name == "ltx-2" and kind != "pipeline": + from .video_ltx2 import is_ltx23_checkpoint + + probe = checkpoint_local + if probe is None: + root = Path(kwargs["repo_id"]).expanduser() + probe = root if root.is_file() else None + ltx23 = probe is not None and is_ltx23_checkpoint(probe) + if ltx23: + expected = self._estimate_download_bytes( + kwargs["repo_id"], + kwargs.get("gguf_filename"), + base, + kwargs.get("hf_token"), + kind, + ltx23 = True, + ) + with self._lock: + if self._load_token == token and self._loading is not None: + self._loading.expected_bytes = expected + base_local = self._predownload_base( + base, kwargs.get("hf_token"), kind, ltx23 = ltx23 + ) + # The 2.3 assembly pulls per component from the hub id (its snapshot here + # deliberately lacks the base VAEs), so it only gets the warmed cache; the + # generic from_pretrained paths get the complete local snapshot. + kwargs["_base_local_dir"] = None if ltx23 else base_local self.load_pipeline(**kwargs) with self._lock: if self._load_token == token: @@ -333,8 +367,13 @@ class VideoBackend: if self._load_token == token and self._loading is not None: self._loading.error = redact_native_paths(str(exc)) + # Base-repo subfolders an LTX-2.3 assembly reads: the checkpoint (plus the GGUF + # repo's extras files) supplies the DiT, connectors, both VAEs and the vocoder, + # so only the 2.0 base's scheduler / text encoder / tokenizer are pulled. + _LTX23_BASE_PREFIXES = ("scheduler/", "text_encoder/", "tokenizer/") + @staticmethod - def _base_download_files(info: Any, kind: str) -> list[tuple[str, int]]: + def _base_download_files(info: Any, kind: str, *, ltx23: bool = False) -> 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, @@ -344,7 +383,10 @@ class VideoBackend: - 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.""" + - ``transformer/`` when a GGUF/single-file checkpoint replaces the DiT; + - everything but scheduler / text encoder / tokenizer for an LTX-2.3 + checkpoint (``ltx23``), whose VAEs/vocoder/connectors come from the + checkpoint and its extras, not the 2.0 base.""" files: list[tuple[str, int]] = [] for sibling in info.siblings or []: name, size = sibling.rfilename, sibling.size or 0 @@ -356,6 +398,12 @@ class VideoBackend: continue if name.startswith("text_encoder/diffusion_pytorch_model"): continue + if ( + ltx23 + and "/" in name + and not name.startswith(VideoBackend._LTX23_BASE_PREFIXES) + ): + continue files.append((name, int(size))) return files @@ -366,6 +414,7 @@ class VideoBackend: base: str, hf_token: Optional[str], kind: str, + ltx23: bool = False, ) -> Optional[int]: """Total bytes this load will pull (checkpoint + companions), or None.""" try: @@ -380,12 +429,14 @@ class VideoBackend: total += int(sibling.size) if base and not Path(base).expanduser().exists(): info = api.model_info(base, files_metadata = True) - total += sum(size for _, size in self._base_download_files(info, kind)) + total += sum(size for _, size in self._base_download_files(info, kind, ltx23 = ltx23)) 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]: + def _predownload_base( + self, base: str, hf_token: Optional[str], kind: str, *, ltx23: bool = False + ) -> 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 @@ -401,7 +452,7 @@ class VideoBackend: 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) + files = self._base_download_files(info, kind, ltx23 = ltx23) 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 diff --git a/studio/backend/tests/test_video_backend.py b/studio/backend/tests/test_video_backend.py index bc15bc0207..3eece7f807 100644 --- a/studio/backend/tests/test_video_backend.py +++ b/studio/backend/tests/test_video_backend.py @@ -509,3 +509,23 @@ def test_pipeline_load_uses_predownloaded_dir(fake_runtime, tmp_path): ) assert _FakePipeline.last["base"] == str(tmp_path) backend.unload() + + +def test_base_download_files_ltx23_keeps_only_shared_components(): + # A 2.3 checkpoint supplies the DiT, connectors, both VAEs and the vocoder, so + # the base pull shrinks to scheduler + text encoder + tokenizer (+ root manifest). + siblings = _LTX2_SIBLINGS + [ + _sibling("scheduler/scheduler_config.json", 1), + _sibling("connectors/diffusion_pytorch_model.safetensors", 3), + _sibling("latent_upsampler/diffusion_pytorch_model.safetensors", 1), + ] + info = types.SimpleNamespace(siblings = siblings) + names = [n for n, _ in VideoBackend._base_download_files(info, "gguf", ltx23 = True)] + assert "model_index.json" in names + assert "scheduler/scheduler_config.json" in names + assert "text_encoder/model-00001-of-00002.safetensors" in names + assert "tokenizer/tokenizer.model" in names + assert not any( + n.startswith(("vae/", "connectors/", "latent_upsampler/", "transformer/")) + for n in names + )