diff --git a/studio/backend/core/inference/diffusion_memory.py b/studio/backend/core/inference/diffusion_memory.py index 300ee82a8d..8706b6d7fd 100644 --- a/studio/backend/core/inference/diffusion_memory.py +++ b/studio/backend/core/inference/diffusion_memory.py @@ -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: diff --git a/studio/backend/core/inference/video.py b/studio/backend/core/inference/video.py index 4bd7751f7c..dcdbb863bb 100644 --- a/studio/backend/core/inference/video.py +++ b/studio/backend/core/inference/video.py @@ -331,8 +331,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( @@ -393,7 +406,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 @@ -463,6 +480,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 diff --git a/studio/backend/core/training/diffusion_dit_trainer.py b/studio/backend/core/training/diffusion_dit_trainer.py index 6d5c068f5d..ba4a4bfcbf 100644 --- a/studio/backend/core/training/diffusion_dit_trainer.py +++ b/studio/backend/core/training/diffusion_dit_trainer.py @@ -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) diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index 74088ab886..35d6921605 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -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, diff --git a/studio/backend/tests/test_video_backend.py b/studio/backend/tests/test_video_backend.py index c34cee5bf9..edb0a41bbb 100644 --- a/studio/backend/tests/test_video_backend.py +++ b/studio/backend/tests/test_video_backend.py @@ -459,6 +459,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), ] @@ -473,7 +474,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(): @@ -528,3 +532,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 == []