From 7a05d8655b8b2e29676b7e7fdbbb62802f8a36f5 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sun, 5 Jul 2026 05:30:58 +0000 Subject: [PATCH] Stream every DiT through group offload, not just the primary transformer A dual-DiT pipeline (Ideogram 4's unconditional tower) placed its second denoiser resident under the group tier, which defeats the tier since the pair rarely fits where one alone did not. Stream transformer_2 and unconditional_transformer alongside the transformer and keep only the smaller companions resident. --- studio/backend/core/inference/diffusion_memory.py | 15 +++++++++++++-- 1 file changed, 13 insertions(+), 2 deletions(-) diff --git a/studio/backend/core/inference/diffusion_memory.py b/studio/backend/core/inference/diffusion_memory.py index c56d1592ce..4c4a3cca2c 100644 --- a/studio/backend/core/inference/diffusion_memory.py +++ b/studio/backend/core/inference/diffusion_memory.py @@ -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: