Merge branch 'video-tab' into video-wan

This commit is contained in:
Daniel Han 2026-07-05 05:36:05 +00:00
commit aa5c9e3a12
5 changed files with 90 additions and 6 deletions

View file

@ -522,6 +522,16 @@ def _apply_group_offload(pipe: Any, device: str, logger: Any) -> bool:
import torch
from diffusers.hooks import apply_group_offloading
# A dual-DiT pipeline (e.g. Ideogram 4's unconditional tower) carries a second
# denoiser as large as the first; leaving it resident would defeat this tier
# (the pair rarely fits where one alone did not). Stream every DiT and keep
# only the genuinely smaller companions resident.
streamed: dict[str, Any] = {"transformer": transformer}
for extra in ("transformer_2", "unconditional_transformer"):
module = getattr(pipe, extra, None)
if isinstance(module, torch.nn.Module):
streamed[extra] = module
onload = torch.device(device)
use_stream = onload.type == "cuda" # overlap H2D copies with compute on CUDA
gkwargs: dict[str, Any] = {
@ -551,11 +561,12 @@ def _apply_group_offload(pipe: Any, device: str, logger: Any) -> bool:
# load-time crash. The streamed transformer manages its own placement via the
# offloading hooks applied next.
for name, comp in getattr(pipe, "components", {}).items():
if name == "transformer":
if name in streamed:
continue
if isinstance(comp, torch.nn.Module):
comp.to(onload)
apply_group_offloading(transformer, **gkwargs)
for module in streamed.values():
apply_group_offloading(module, **gkwargs)
return True
except Exception as exc: # noqa: BLE001 — fall back to whole-module offload
if logger is not None:

View file

@ -428,8 +428,21 @@ class VideoBackend:
probe = checkpoint_local
if probe is None:
# Local repos: a bare file, or a directory whose child the same
# resolver load_pipeline uses picks out. Unresolvable here means
# load_pipeline will surface the real error; keep the wide pull.
root = Path(kwargs["repo_id"]).expanduser()
probe = root if root.is_file() else None
if root.is_file():
probe = root
elif root.is_dir():
try:
probe = self._resolve_checkpoint_path(
kwargs["repo_id"],
kwargs.get("gguf_filename"),
kwargs.get("hf_token"),
)
except Exception: # noqa: BLE001 -- surfaced by load_pipeline
probe = None
ltx23 = probe is not None and is_ltx23_checkpoint(probe)
if ltx23:
expected = self._estimate_download_bytes(
@ -490,7 +503,11 @@ class VideoBackend:
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")):
# .jinja: tokenizer/chat_template.jinja ships as a standalone file in the
# LTX-2 and HunyuanVideo-1.5 repos (not embedded in tokenizer_config.json)
# and apply_chat_template needs it at generation time, so a snapshot
# without it loads fine and then crashes the first generation.
if not name.endswith((".safetensors", ".json", ".model", ".txt", ".jinja")):
continue
if "/" not in name and name.endswith(".safetensors"):
continue
@ -560,6 +577,11 @@ class VideoBackend:
snapshot_root: Optional[Path] = None
for name, _ in files:
# Explicit per-file check: a fully-cached file returns without ever
# consulting the event, so a warm-cache sweep would otherwise run to
# completion after an unload already cancelled this load.
if self._cancel_event.is_set():
raise RuntimeError(VIDEO_CANCELLED_MSG)
local = Path(
hf_hub_download_with_xet_fallback(
base, name, hf_token, cancel_event = self._cancel_event

View file

@ -1173,6 +1173,13 @@ def run_dit_lora_training(
device = "cuda" if torch.cuda.is_available() else "cpu"
# The flow-matching + 4-bit path is bf16 throughout (fp32 on a CPU-only box, which is
# unsupported for real runs but keeps import/unit tests architecture-agnostic).
# Fail fast on pre-Ampere CUDA (T4/V100/RTX 20xx): bf16 compute is required and the run
# would otherwise die deep in model load with an opaque dtype error.
if device == "cuda" and not torch.cuda.is_bf16_supported():
raise ValueError(
"This trainer requires a bfloat16-capable GPU (Ampere or newer); "
"this CUDA device does not support bf16."
)
weight_dtype = torch.bfloat16 if device == "cuda" else torch.float32
_assert_trusted_base_model(cfg.base_model)

View file

@ -1339,7 +1339,12 @@ async def get_diffusion_training_run(
rec = get_diffusion_run(job_id)
if rec is None:
raise HTTPException(status_code = 404, detail = "No such training run.")
return DiffusionTrainingRunDetail(**rec)
try:
return DiffusionTrainingRunDetail(**rec)
except ValidationError:
# A malformed on-disk record (hand-edited / older shape) should read as absent
# rather than 500 the endpoint, mirroring how the list route skips bad records.
raise HTTPException(status_code = 404, detail = "No such training run.")
# Extensions accepted into an image-training dataset folder: images the trainer reads,

View file

@ -889,6 +889,7 @@ _LTX2_SIBLINGS = [
_sibling("text_encoder/diffusion_pytorch_model-00002-of-00002.safetensors", 25),
_sibling("vae/diffusion_pytorch_model.safetensors", 3),
_sibling("tokenizer/tokenizer.model", 1),
_sibling("tokenizer/chat_template.jinja", 1),
_sibling("assets/example.mp4", 500),
]
@ -903,7 +904,10 @@ def test_base_download_files_scopes_pipeline_pull():
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
# The standalone chat template must survive the whitelist: apply_chat_template
# reads it at generation time and it is not embedded in tokenizer_config.json.
assert "tokenizer/chat_template.jinja" in files
assert sum(files.values()) == 10 + 1 + 20 + 18 + 25 + 25 + 3 + 1 + 1
def test_base_download_files_gguf_drops_transformer():
@ -958,3 +962,38 @@ def test_base_download_files_ltx23_keeps_only_shared_components():
assert not any(
n.startswith(("vae/", "connectors/", "latent_upsampler/", "transformer/")) for n in names
)
def test_predownload_base_honors_cancel_between_files(monkeypatch):
# A warm-cache sweep returns each file instantly without consulting the event,
# so the loop must check it explicitly or an unload mid-predownload is ignored.
backend = VideoBackend()
backend._cancel_event.set()
calls: list = []
monkeypatch.setattr(
"utils.hf_xet_fallback.hf_hub_download_with_xet_fallback",
lambda repo, fn, tok, **kw: (calls.append(fn), f"/cache/{fn}")[1],
)
class _Api:
def __init__(self, token = None):
pass
def model_info(
self,
repo,
files_metadata = True,
):
return types.SimpleNamespace(
siblings = [
_sibling("model_index.json", 1),
_sibling("vae/diffusion_pytorch_model.safetensors", 2),
]
)
import huggingface_hub
monkeypatch.setattr(huggingface_hub, "HfApi", _Api)
with pytest.raises(RuntimeError, match = "cancelled"):
backend._predownload_base("base/repo", None, "pipeline")
assert calls == []