Merge branch 'video-tab' into video-wan
This commit is contained in:
commit
aa5c9e3a12
5 changed files with 90 additions and 6 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue