"""dexp012_blob_supervised.py — exp012: SEGMENTATION-BLOB SUPERVISION on the HIGH-noise band (Phil's blobbing role, first real qualitatively-different supervision; day-3 plan item 2 — design preregistered there). Coupling (conservative, controlled): foreground-weighted LOW-PASS x0 loss on the HIGH band only: x0_hat = (x_t - sqrt(1-abar_t) * eps_hat) / sqrt(abar_t) L_blob = mean( blob ⊙ (LP(x0_hat) - LP(x0))^2 ), lambda = 0.5 routed by the crossfade windows (w2 only). Differs from exp009's inert reweighting: x0-space structural prediction at high noise + EXTERNAL segmentation information (the rebuilt 100%-non-empty blob targets). In-bed gauges for ALL arms (blob-mb3 trained here; uniform-mb3 + mono48 loaded from exp011 ckpts; frozen): common eps gauge per band + BLOB GAUGE (foreground LP-x0 error) per band. Prereg P1-P4 in the day-3 plan. Pod: bash pod2/run_exp012.sh [DEXP12_SEED=0 DEXP12_LAMBDA=0.5] """ from __future__ import annotations import json import os import sys import time sys.path[:0] = ["pod2", "."] import torch import torch.nn.functional as F from pod_ledger import ledger_run, note, burn_down from d1_substrate import MEM_FRACTION from dexp006_sd15core_relay import make_schedule, add_noise from dexp008_multiband import (MultibandDelta, MonoDelta, band_weights, band_of, attach, load_unet, N_BANDS) from dexp009_bandroles import lp N_TRAIN, N_VAL = 4096, 256 BATCH = int(os.environ.get("DEXP12_BATCH", "16")) STEPS = int(os.environ.get("DEXP12_STEPS", "3000")) SEED = int(os.environ.get("DEXP12_SEED", "0")) LAM = float(os.environ.get("DEXP12_LAMBDA", "0.5")) LR, CFG_DROPOUT = 1e-3, 0.1 D11 = ("/workspace/data/dexp011" if os.path.isdir("/workspace") else "./data/dexp011") D11_CKPT = ("/workspace/ckpts2/dexp011" if os.path.isdir("/workspace") else D11) DATA_DIR = ("/workspace/data/dexp012" if os.path.isdir("/workspace") else os.path.join(os.environ.get("GEOLIP_DATA", "./data"), "dexp012")) CKPT_DIR = ("/workspace/ckpts2/dexp012" if os.path.isdir("/workspace") else DATA_DIR) def x0_from_eps(noisy, eps, t, acp): a = acp[t].sqrt()[:, None, None, None] s = (1 - acp[t]).sqrt()[:, None, None, None] return (noisy - s * eps) / a.clamp_min(1e-4) def blob_lp_err(x0_hat, x0, blob): """Foreground-weighted LP-x0 squared error, per sample (B,).""" d2 = (lp(x0_hat) - lp(x0)) ** 2 # (B,4,64,64) m = blob[:, None] # (B,1,64,64) denom = m.sum(dim=(1, 2, 3)).clamp_min(1.0) * d2.shape[1] return (d2 * m).sum(dim=(1, 2, 3)) / denom def run(device="cuda"): torch.cuda.set_per_process_memory_fraction(MEM_FRACTION, 0) os.makedirs(CKPT_DIR, exist_ok=True) os.makedirs(DATA_DIR, exist_ok=True) acp = make_schedule(device) cache = torch.load(os.path.join(D11, "cache.pt"), map_location="cpu", weights_only=True) assert int(cache["blob_empty_count"]) == 0, \ "blob targets not the rebuilt ones — run fix_blob_targets first" def set_w(wraps, s01): w = band_weights(s01) for wr in wraps: wr.w_bands = w return w def loss_of(unet, wraps, lat, ehs, blob, gen): bsz = lat.shape[0] drop = torch.rand(bsz, generator=gen, device=device) < CFG_DROPOUT ehs = ehs.clone() ehs[drop] = 0 t = torch.randint(0, 1000, (bsz,), generator=gen, device=device) w = set_w(wraps, t.float() / 1000.0) noise = torch.randn(lat.shape, generator=gen, device=device) noisy = add_noise(lat, noise, t, acp) pred = unet(noisy, t, ehs, return_dict=False)[0] base = ((pred - noise) ** 2).mean(dim=(1, 2, 3)) x0h = x0_from_eps(noisy, pred, t, acp) blob_term = blob_lp_err(x0h, lat, blob) return (base + LAM * w[:, 2] * blob_term).mean() @torch.no_grad() def val(unet, wraps): """Common + blob gauges per band, identical math for every arm.""" per_band = {0: [], 1: [], 2: []} blob_band = {0: [], 1: [], 2: []} tot = [] for i in range(0, N_VAL, 32): lat = cache["val_lat"][i:i + 32].to(device) ehs = cache["val_ehs"][i:i + 32].to(device) noise = cache["val_noise"][i:i + 32].to(device) t = cache["val_t"][i:i + 32].to(device) blob = cache["val_blob"][i:i + 32].float().to(device) set_w(wraps, t.float() / 1000.0) noisy = add_noise(lat, noise, t, acp) pred = unet(noisy, t, ehs, return_dict=False)[0] mse = ((pred - noise) ** 2).mean(dim=(1, 2, 3)) bg = blob_lp_err(x0_from_eps(noisy, pred, t, acp), lat, blob) tot += mse.tolist() for j, tv in enumerate((t.float() / 1000.0).tolist()): b = band_of(tv) per_band[b].append(mse[j].item()) blob_band[b].append(bg[j].item()) return (sum(tot) / len(tot), {f"band{b}": round(sum(v) / max(len(v), 1), 6) for b, v in per_band.items()}, {f"band{b}": round(sum(v) / max(len(v), 1), 6) for b, v in blob_band.items()}) results = {"config": {"lambda": LAM, "steps": STEPS, "batch": BATCH, "seed": SEED, "coupling": "foreground-weighted LP-x0, HIGH band only"}} # references gauged IN-BED (the exp009 instrument lesson) with ledger_run(f"dexp012 reference gauges s{SEED}", budget_h=0.4) as h: unet = load_unet(device) v, pb, bb = val(unet, []) results["frozen"] = {"val": v, "per_band": pb, "blob_gauge": bb} del unet torch.cuda.empty_cache() for label, mk, fn in (("uniform_mb3", lambda d: MultibandDelta(d), "multiband3_s0.pt"), ("monolith", lambda d: MonoDelta(d), "monolith_s0.pt")): unet = load_unet(device) mods, wraps = attach(unet, mk) ck = torch.load(os.path.join(D11_CKPT, fn), map_location="cpu", weights_only=True) for m, sd in zip(mods, ck["mods"]): m.load_state_dict(sd, strict=True) v, pb, bb = val(unet, wraps) results[label] = {"val": v, "per_band": pb, "blob_gauge": bb} del unet, mods torch.cuda.empty_cache() h["verdict"] = "refs gauged" with ledger_run(f"dexp012 blob-mb3 s{SEED}", budget_h=2.5) as h: unet = load_unet(device) mods, wraps = attach(unet, lambda d: MultibandDelta(d)) for m in mods: m.assert_zero_init() opt = torch.optim.Adam(mods.parameters(), lr=LR, weight_decay=0.0) gen = torch.Generator(device=device).manual_seed(SEED + 42) idx = torch.Generator().manual_seed(SEED + 7) t0 = time.time() for step in range(1, STEPS + 1): sel = torch.randint(0, N_TRAIN, (BATCH,), generator=idx) loss = loss_of(unet, wraps, cache["lat"][sel].to(device), cache["ehs"][sel].to(device), cache["blob"][sel].float().to(device), gen) loss.backward() opt.step() opt.zero_grad(set_to_none=True) if step == 50 or step % 500 == 0: print(f"[blob] step {step}: loss {loss.item():.4f} | " f"{(time.time() - t0) / step:.2f}s/step", flush=True) v, pb, bb = val(unet, wraps) for m in mods: m.enabled = False v_off, _, _ = val(unet, wraps) d = abs(v_off - results["frozen"]["val"]) assert d < 1e-9, f"toggle parity broken: {d}" for m in mods: m.enabled = True torch.save({"mods": [m.state_dict() for m in mods]}, os.path.join(CKPT_DIR, f"blob_mb3_s{SEED}.pt")) results["blob_mb3"] = {"val": v, "per_band": pb, "blob_gauge": bb, "val_toggled_off": v_off} del unet, mods torch.cuda.empty_cache() h["verdict"] = f"val {v:.5f} blobH {bb['band2']}" bm, um = results["blob_mb3"], results["uniform_mb3"] results["verdict"] = { "P1_blob_gauge_high_band": {"blob_mb3": bm["blob_gauge"]["band2"], "uniform_mb3": um["blob_gauge"]["band2"], "hit": bm["blob_gauge"]["band2"] < um["blob_gauge"]["band2"]}, "P2_common_undegraded": abs(bm["val"] - um["val"]) / um["val"] < 0.005 or bm["val"] < um["val"], "common_val": {"blob": bm["val"], "uniform": um["val"], "mono": results["monolith"]["val"]}, "note": "P4 (image battery on the blob ckpt) runs via dexp010 env; " "1-seed CANDIDATE", } with open(os.path.join(DATA_DIR, "results.json" if SEED == 0 else f"results_s{SEED}.json"), "w") as f: json.dump(results, f, indent=2) note(f"dexp012: {json.dumps(results['verdict']['P1_blob_gauge_high_band'])}") print(json.dumps(results["verdict"], indent=2)) burn_down() return results def smoke(): acp = torch.linspace(0.9999, 0.05, 1000) x = torch.randn(2, 4, 64, 64) n = torch.randn_like(x) t = torch.tensor([100, 900]) a = acp[t].sqrt()[:, None, None, None] s = (1 - acp[t]).sqrt()[:, None, None, None] noisy = a * x + s * n x0h = x0_from_eps(noisy, n, t, acp) assert (x0h - x).abs().max() < 1e-4, "x0 inversion wrong" blob = torch.zeros(2, 64, 64) blob[0, 10:30, 10:30] = 1 e = blob_lp_err(x0h, x, blob) assert e.shape == (2,) and e[0] < 1e-6 and e[1] == 0, e print("dexp012 smoke PASSED (x0 inversion exact + fg-weighted gauge)") if __name__ == "__main__": if "--run" in sys.argv: run() else: smoke()