Merge branch 'diffusion-train-perf' into diffusion-train-precision

This commit is contained in:
Daniel Han 2026-07-05 05:35:52 +00:00
commit 1919491665
2 changed files with 20 additions and 2 deletions

View file

@ -502,6 +502,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] = {
@ -531,11 +541,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

@ -983,6 +983,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)