"""GPU batch-1 pre-flight for SpatialDiffuseSlot (before real launch / new server). Loads the EXACT launch config, runs 3 real optimizer steps at batch=1 on ONE GPU: - tok_L init OK, forward losses finite, backward+step OK - frozen trunk stays frozen (grad None), memory footprint printed - EMA-style second forward after step (params updated, still finite) """ import torch, time from omegaconf import OmegaConf from semanticist.engine.trainer_utils import instantiate_from_config cfg = OmegaConf.load("configs/tokenizer_l_spatial.yaml") mp = cfg.trainer.params.model model = instantiate_from_config(mp).cuda() model.train() opt = torch.optim.AdamW([p for p in model.parameters() if p.requires_grad], lr=1e-4, betas=(0.9, 0.95)) print(f"[pre] trainable={sum(p.numel() for p in model.parameters() if p.requires_grad)/1e6:.1f}M " f"total={sum(p.numel() for p in model.parameters())/1e6:.1f}M") losses_log = [] for step in range(3): x = torch.randn(1, 3, 256, 256, device="cuda") # batch=1 t0 = time.time() with torch.autocast("cuda", dtype=torch.bfloat16): losses = model(x, sample=False) loss = sum(losses.values()) opt.zero_grad(set_to_none=True) loss.backward() # frozen trunk check bad = [n for n, p in model.dit.named_parameters() if n.startswith("blocks.") and p.grad is not None] assert not bad, f"frozen trunk got grads: {bad[:3]}" torch.nn.utils.clip_grad_norm_( [p for p in model.parameters() if p.requires_grad], 3.0) opt.step() dt = time.time() - t0 l = {k: float(v) for k, v in losses.items()} losses_log.append(l) print(f"[step {step}] {l} | {dt:.2f}s | mem {torch.cuda.max_memory_allocated()/2**30:.1f}GiB") assert torch.isfinite(loss) print("GPU BATCH-1 PRE-FLIGHT PASSED ✅ (3 steps, finite, frozen trunk intact)")