exp012 shipped: blob v1 honest negative BUT the role-aligned gauge reveals multiband beats monolith ~10% on HIGH-band foreground structure; ANIMA R0b all green (dtype violation caught+fixed)
4e7b91c verified | """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() | |
| 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() | |