Merge branch 'diffusion-train-perf' into diffusion-train-precision
This commit is contained in:
commit
1919491665
2 changed files with 20 additions and 2 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue