diff --git a/studio/backend/core/training/diffusion_dit_trainer.py b/studio/backend/core/training/diffusion_dit_trainer.py index 4f693e9f4c..fdc7fa96b2 100644 --- a/studio/backend/core/training/diffusion_dit_trainer.py +++ b/studio/backend/core/training/diffusion_dit_trainer.py @@ -40,10 +40,13 @@ from core.training.diffusion_train_common import ( DEFAULT_LORA_TARGETS, DiffusionLoraConfig, EventCb, + LATENT_CACHE_OVER_BUDGET, StopCb, _apply_perf_flags, _assert_trusted_base_model, _emit, + _latent_cache_forced, + _latent_cache_over_budget, _plan_cache_variants, _publish_to_lora_catalog, _restore_perf_flags, @@ -826,6 +829,9 @@ def _build_latent_cache(spec, vae, image_paths, cfg, device, weight_dtype, on_ev cache: list[list[tuple]] = [] total = len(image_paths) + total_variants = sum(len(v) for v in plan) + forced = _latent_cache_forced() + gated = False for i, path in enumerate(image_paths): variants = [] for u_left, u_top, flip in plan[i]: @@ -837,7 +843,30 @@ def _build_latent_cache(spec, vae, image_paths, cfg, device, weight_dtype, on_ev .to(device) ) a, b = spec.encode_latent_stats(vae, px) - variants.append((_hold(a), _hold(b))) + a, b = _hold(a), _hold(b) + if not forced and not gated: + # Size-gate the automatic cache off the first REAL encoded variant, before + # building the rest: packed 16-channel DiT latents x variants x images of two + # fp32 tensors can exhaust host/pinned RAM. Over budget we bail with the VAE + # still resident so the loop encodes latents per step instead. ``b`` is None + # for a deterministic-latent family, so only ``a`` contributes bytes there. + per_variant = a.numel() * a.element_size() + if b is not None: + per_variant += b.numel() * b.element_size() + if _latent_cache_over_budget(per_variant, total_variants): + _emit( + on_event, + "warning", + message = ( + "Latent cache disabled: estimated " + f"{per_variant * total_variants / 1024 ** 3:.1f} GiB over the " + "budget; encoding latents per step instead. Set " + "UNSLOTH_DIFFUSION_FORCE_LATENT_CACHE=1 to keep it." + ), + ) + return LATENT_CACHE_OVER_BUDGET + gated = True + variants.append((a, b)) cache.append(variants) if (i + 1) % 4 == 0 or i + 1 == total: _emit(on_event, "preparing", stage = "cache_latents", done = i + 1, total = total) @@ -1064,7 +1093,11 @@ def _train_dit(cfg, spec, pairs, rng, device, weight_dtype, on_event, _check_sto latent_cache = _build_latent_cache( spec, vae, image_paths, cfg, device, weight_dtype, on_event, _check_stop ) - if latent_cache is None: # stopped during the cache build; nothing trained yet + if latent_cache is LATENT_CACHE_OVER_BUDGET: + # The estimated cache exceeded the host-memory budget; keep the VAE resident and + # fall through to the in-loop encode path (latent_cache stays None). + latent_cache = None + elif latent_cache is None: # stopped during the cache build; nothing trained yet _emit( on_event, "complete", @@ -1074,15 +1107,16 @@ def _train_dit(cfg, spec, pairs, rng, device, weight_dtype, on_event, _check_sto steps_run = 0, ) return str(out_dir) - try: - pipe.vae = None - except Exception: # noqa: BLE001 -- a pipeline without a settable vae keeps it - pass - del vae - vae = None - gc.collect() - if device == "cuda": - torch.cuda.empty_cache() + else: + try: + pipe.vae = None + except Exception: # noqa: BLE001 -- a pipeline without a settable vae keeps it + pass + del vae + vae = None + gc.collect() + if device == "cuda": + torch.cuda.empty_cache() # Variant picks use their own stream so the training loop's index/noise draws stay on # the same seed-deterministic sequence whether or not the cache is enabled. variant_rng = random.Random(cfg.seed + 1) diff --git a/studio/backend/core/training/diffusion_lora_trainer.py b/studio/backend/core/training/diffusion_lora_trainer.py index b7c65a0ac4..8c3a79afaf 100644 --- a/studio/backend/core/training/diffusion_lora_trainer.py +++ b/studio/backend/core/training/diffusion_lora_trainer.py @@ -47,12 +47,15 @@ from core.training.diffusion_train_common import ( # noqa: F401 EventCb, StopCb, DiffusionLoraConfig, + LATENT_CACHE_OVER_BUDGET, _apply_perf_flags, _assert_trusted_base_model, _coerce_gradient_checkpointing, _config_from_dict, _CONFIG_ALIASES, _emit, + _latent_cache_forced, + _latent_cache_over_budget, _plan_cache_variants, _publish_to_lora_catalog, _restore_perf_flags, @@ -202,6 +205,9 @@ def _build_sdxl_latent_cache( cache: list[list[tuple]] = [] total = len(image_paths) + total_variants = sum(len(v) for v in plan) + forced = _latent_cache_forced() + gated = False for i, path in enumerate(image_paths): variants = [] for u_left, u_top, flip in plan[i]: @@ -213,6 +219,25 @@ def _build_sdxl_latent_cache( dist = vae.encode(pixel_values).latent_dist a = _hold(dist.mean * vae_scale) b = _hold(dist.std * vae_scale) + if not forced and not gated: + # Size-gate the automatic cache off the first REAL encoded variant, before + # building the rest: thousands of images x variants of two fp32 tensors can + # exhaust host/pinned RAM with no fallback. Over budget we bail with the VAE + # still resident so the loop encodes latents per step instead. + per_variant = a.numel() * a.element_size() + b.numel() * b.element_size() + if _latent_cache_over_budget(per_variant, total_variants): + _emit( + on_event, + "warning", + message = ( + "Latent cache disabled: estimated " + f"{per_variant * total_variants / 1024 ** 3:.1f} GiB over the " + "budget; encoding latents per step instead. Set " + "UNSLOTH_DIFFUSION_FORCE_LATENT_CACHE=1 to keep it." + ), + ) + return LATENT_CACHE_OVER_BUDGET + gated = True variants.append((a, b, tuple(time_ids))) cache.append(variants) if (i + 1) % 4 == 0 or i + 1 == total: @@ -404,7 +429,11 @@ def run_diffusion_lora_training( on_event, _check_stop, ) - if latent_cache is None: # stopped during the cache build; nothing trained yet + if latent_cache is LATENT_CACHE_OVER_BUDGET: + # The estimated cache exceeded the host-memory budget; keep the VAE resident + # and fall through to the in-loop encode path (latent_cache stays None). + latent_cache = None + elif latent_cache is None: # stopped during the cache build; nothing trained yet out_dir = Path(cfg.output_dir).expanduser() _emit( on_event, @@ -415,15 +444,16 @@ def run_diffusion_lora_training( steps_run = 0, ) return str(out_dir) - try: - pipe.vae = None - except Exception: # noqa: BLE001 -- a pipeline without a settable vae keeps it - pass - del vae - vae = None - gc.collect() - if device == "cuda": - torch.cuda.empty_cache() + else: + try: + pipe.vae = None + except Exception: # noqa: BLE001 -- a pipeline without a settable vae keeps it + pass + del vae + vae = None + gc.collect() + if device == "cuda": + torch.cuda.empty_cache() # Variant picks use their own stream so the loop's index/noise draws stay on the same # seed-deterministic sequence whether or not the cache is enabled. variant_rng = random.Random(cfg.seed + 1) diff --git a/studio/backend/core/training/diffusion_train_common.py b/studio/backend/core/training/diffusion_train_common.py index 220f4cff32..2f5d063e69 100644 --- a/studio/backend/core/training/diffusion_train_common.py +++ b/studio/backend/core/training/diffusion_train_common.py @@ -517,6 +517,42 @@ def _plan_cache_variants( return plan +# Host-memory budget for the AUTOMATIC latent cache. The cache holds two fp32 posterior +# tensors (mean/std, VAE scale folded in) per crop/flip variant per image, pinned on a CUDA +# host. At 1024px an SDXL variant is ~0.5 MiB and a 16-channel DiT variant several times +# that, so a few thousand images x cache_variants can exhaust host or pinned RAM with no +# fallback. Over this budget the default falls back to per-step VAE encoding. A fixed +# constant (rather than a psutil RAM fraction) keeps the gate dependency-free and identical +# across hosts; it is deliberately conservative, well under a typical training host's RAM. +_LATENT_CACHE_BUDGET_BYTES = 4 * 1024 ** 3 # 4 GiB + +# Returned by the cache builders when the estimated cache exceeds the budget: the caller +# keeps the VAE resident and encodes each step's latents in-loop. A distinct sentinel from +# ``None`` (which means a stop was requested mid-build) so the two are not conflated. +LATENT_CACHE_OVER_BUDGET: Any = object() + + +def _latent_cache_forced() -> bool: + """The user explicitly forced the latent cache on, bypassing the size gate. This is the + explicit opt-in counterpart to ``UNSLOTH_DIFFUSION_NO_LATENT_CACHE`` (the explicit + opt-out); only the automatic default is size-gated, so an explicit choice is honoured + verbatim in either direction.""" + return os.environ.get("UNSLOTH_DIFFUSION_FORCE_LATENT_CACHE", "") in ("1", "true") + + +def _latent_cache_over_budget( + per_variant_bytes: int, total_variants: int, budget_bytes: Optional[int] = None +) -> bool: + """True when a cache of ``total_variants`` entries, each two fp32 tensors totalling + ``per_variant_bytes``, is estimated to exceed ``budget_bytes``. ``per_variant_bytes`` is + measured from a real encoded latent, so the estimate tracks the actual per-family tensor + shape (SDXL 4-channel vs. a packed 16-channel DiT latent) rather than a guess. The budget + is read from the module constant at call time when not given, so tests can override it.""" + if budget_bytes is None: + budget_bytes = _LATENT_CACHE_BUDGET_BYTES + return per_variant_bytes * max(0, total_variants) > budget_bytes + + def _apply_perf_flags( cfg: "DiffusionLoraConfig", device: str, diff --git a/studio/backend/tests/test_diffusion_train_perf.py b/studio/backend/tests/test_diffusion_train_perf.py index c0188a9694..162822ec57 100644 --- a/studio/backend/tests/test_diffusion_train_perf.py +++ b/studio/backend/tests/test_diffusion_train_perf.py @@ -32,11 +32,16 @@ from core.training.diffusion_dit_trainer import ( ) from core.training.diffusion_train_common import ( DiffusionLoraConfig, + LATENT_CACHE_OVER_BUDGET, _apply_perf_flags, _config_from_dict, + _latent_cache_forced, + _latent_cache_over_budget, _plan_cache_variants, _restore_perf_flags, ) +import core.training.diffusion_lora_trainer as sdxl_trainer +import core.training.diffusion_train_common as train_common from core.training.diffusion_training_service import DiffusionTrainingService from models.training import DiffusionTrainingStartRequest, DiffusionTrainingStopRequest from routes.training import router as training_router @@ -375,3 +380,90 @@ def test_perf_flags_tf32_off_clears_flags(): torch.get_float32_matmul_precision(), ) assert after == before + + +# ── latent cache size gate ──────────────────────────────────────────────────── +class _FakeLatentDist: + def __init__(self, shape): + self.mean = torch.zeros(shape, dtype = torch.float32) + self.std = torch.ones(shape, dtype = torch.float32) + + +class _FakeEncoded: + def __init__(self, shape): + self.latent_dist = _FakeLatentDist(shape) + + +class _FakeVae: + # Minimal VAE stand-in: encode() returns a posterior of the requested latent shape so the + # builder measures a real per-variant byte size without a model load or image files. + def __init__(self, shape): + self._shape = shape + + def encode(self, pixel_values): + return _FakeEncoded(self._shape) + + +def _fake_planned_loader(path, resolution, center_crop, u_left, u_top, flip): + # The fake VAE ignores pixels; return a valid tensor + square SDXL time_ids. + tensor = torch.zeros(3, resolution, resolution, dtype = torch.float32) + return tensor, (resolution, resolution, 0, 0, resolution, resolution) + + +def _build_fake_sdxl_cache(monkeypatch, num_images, latent_shape): + # center_crop + no flip collapses to one variant per image, so total_variants == num_images. + monkeypatch.setattr(sdxl_trainer, "_load_image_tensor_planned", _fake_planned_loader) + cfg = _cfg(cache_variants = 1, center_crop = True, random_flip = False).normalized() + return sdxl_trainer._build_sdxl_latent_cache( + _FakeVae(latent_shape), + 1.0, + [f"img{i}.png" for i in range(num_images)], + cfg, + "cpu", + torch.float32, + None, + lambda: False, + ) + + +def test_latent_cache_over_budget_boundary(): + # 32 bytes per variant x 4 variants = 128 bytes; exactly at budget is not "over". + assert _latent_cache_over_budget(32, 4, budget_bytes = 200) is False + assert _latent_cache_over_budget(32, 4, budget_bytes = 128) is False + assert _latent_cache_over_budget(32, 4, budget_bytes = 127) is True + # An empty plan can never overflow. + assert _latent_cache_over_budget(1_000_000, 0, budget_bytes = 1) is False + + +def test_latent_cache_forced_env(monkeypatch): + monkeypatch.delenv("UNSLOTH_DIFFUSION_FORCE_LATENT_CACHE", raising = False) + assert _latent_cache_forced() is False + monkeypatch.setenv("UNSLOTH_DIFFUSION_FORCE_LATENT_CACHE", "1") + assert _latent_cache_forced() is True + + +def test_sdxl_cache_built_under_budget(monkeypatch): + # Default (4 GiB) budget: a handful of tiny latents fits, so the full cache is returned. + monkeypatch.delenv("UNSLOTH_DIFFUSION_FORCE_LATENT_CACHE", raising = False) + cache = _build_fake_sdxl_cache(monkeypatch, num_images = 3, latent_shape = (1, 4, 8, 8)) + assert cache is not LATENT_CACHE_OVER_BUDGET and cache is not None + assert len(cache) == 3 + assert all(len(variants) == 1 for variants in cache) + + +def test_sdxl_cache_gated_over_budget(monkeypatch): + # A budget below one variant forces the gate on the first encode: the sentinel is returned + # so the caller keeps the VAE resident and encodes per step. + monkeypatch.delenv("UNSLOTH_DIFFUSION_FORCE_LATENT_CACHE", raising = False) + monkeypatch.setattr(train_common, "_LATENT_CACHE_BUDGET_BYTES", 8) + cache = _build_fake_sdxl_cache(monkeypatch, num_images = 3, latent_shape = (1, 4, 8, 8)) + assert cache is LATENT_CACHE_OVER_BUDGET + + +def test_sdxl_cache_force_bypasses_gate(monkeypatch): + # An explicit force-on must be honoured verbatim even when the estimate is over budget. + monkeypatch.setenv("UNSLOTH_DIFFUSION_FORCE_LATENT_CACHE", "1") + monkeypatch.setattr(train_common, "_LATENT_CACHE_BUDGET_BYTES", 8) + cache = _build_fake_sdxl_cache(monkeypatch, num_images = 3, latent_shape = (1, 4, 8, 8)) + assert cache is not LATENT_CACHE_OVER_BUDGET and cache is not None + assert len(cache) == 3