| """FULL SMOKE for SpatialDiffuseSlot: |
| 1. tok_L init loads: encoder 100%, DiT trunk (except null_cond K-mismatch) |
| 2. training forward: diff_loss + repa_loss finite |
| 3. backward: pool/cond-embedder grads flow; frozen DiT trunk has no grads |
| 4. mask actually in play (spatial_mask buffer on the rebuilt DiT) |
| CPU, small batch. Catches integration bugs before GPU run. |
| """ |
| import torch |
| from omegaconf import OmegaConf |
| from semanticist.stage1.spatial_diffuse_slot import SpatialDiffuseSlot |
|
|
| torch.manual_seed(0) |
| cfg = OmegaConf.load("configs/tokenizer_l.yaml") |
| p = dict(cfg["trainer"]["params"]["model"]["params"]) |
| p.pop("ckpt_path", None) |
| |
| m = SpatialDiffuseSlot(freeze_dit=True, pool_depth=2, **p) |
|
|
| ck = torch.load("../new_eval_spatial_reasoning0430/eval_assets/semanticist_pretrained/semanticist_tok_L.pkl", map_location="cpu") |
| ck = {k.replace("._orig_mod", ""): v for k, v in ck.items()} |
| |
| |
| ck.pop("dit.null_cond", None) |
| ret = m.load_state_dict(ck, strict=False) |
| enc_missing = [k for k in ret.missing_keys if k.startswith("encoder.")] |
| dit_missing = [k for k in ret.missing_keys if k.startswith("dit.") and "null_cond" not in k and "pos_embed" not in k] |
| print(f"[init] encoder missing={len(enc_missing)} (0 expected) | dit trunk missing={len(dit_missing)} (0 expected)") |
| print(f"[init] null_cond shape={tuple(m.dit.null_cond.shape)} (1,85,16 expected)") |
| assert len(enc_missing) == 0 and len(dit_missing) == 0 |
| assert tuple(m.dit.null_cond.shape) == (1, 85, 16) |
| assert m.dit.spatial_mask.shape == (64 + 85, 64 + 85) |
|
|
| |
| m.train() |
| x = torch.randn(2, 3, 256, 256) |
| losses = m(x, sample=False) |
| print("[fwd] losses:", {k: float(v) for k, v in losses.items()}) |
| total = sum(losses.values()) |
| assert torch.isfinite(total) |
|
|
| |
| total.backward() |
| def gnorm(mod): |
| s = 0.0 |
| for q in mod.parameters(): |
| if q.grad is not None: |
| s += q.grad.norm().item() ** 2 |
| return s ** 0.5 |
| g_pool = gnorm(m.spatial_pool) |
| g_emb = gnorm(m.dit.autoenc_cond_embedder) |
| g_trunk = sum((q.grad is not None and q.grad.abs().sum().item() > 0) for n, q in m.dit.named_parameters() if n.startswith("blocks.")) |
| print(f"[bwd] grad-norm pool={g_pool:.4f} cond_embedder={g_emb:.4f} | frozen-trunk grads={g_trunk} (0 expected)") |
| assert g_pool > 0 and g_emb > 0 and g_trunk == 0 |
| print("FULL SMOKE PASSED ✅ init/forward/backward/freeze all correct") |
|
|