Intrinsic Schrödinger Bridge Diffusion / stage2_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json
  2import random
  3import sys
  4import numpy as np
  5import torch
  6from torch import nn
  7
  8sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
  9from bench import make_model, train_model, evaluate, sweep_baseline, make_report
 10from manifold_dynamics_track import get_dataset
 11
 12SEEDS = tuple(range(8))
 13NTR, NTE = 400, 200
 14EPOCHS = 18
 15BATCH = 128
 16LRS = [1e-3, 3e-3, 1e-2]
 17
 18
 19def seed_all(seed):
 20    random.seed(seed)
 21    np.random.seed(seed)
 22    torch.manual_seed(seed)
 23    if torch.cuda.is_available():
 24        torch.cuda.manual_seed_all(seed)
 25
 26
 27def dataset(seed):
 28    d = get_dataset(seed, NTR, NTE)
 29    for k in ("xtr", "ytr", "xte", "yte"):
 30        d[k] = torch.as_tensor(d[k], dtype=torch.float32)
 31    d["input_shape"] = tuple(d["xtr"].shape[1:])
 32    d["out_dim"] = 3
 33    return d
 34
 35
 36def baseline_factory(cfg):
 37    def run(seed):
 38        seed_all(seed)
 39        d = dataset(seed)
 40        net = make_model("rnn_small", d["input_shape"], 3)
 41        _, metric, _ = train_model(net, d, epochs=EPOCHS, lr=cfg["lr"], batch=BATCH,
 42                                   log=lambda *_: None)
 43        return float(metric)
 44    return run
 45
 46
 47def intrinsic_factory(cfg, collect=False):
 48    def run(seed):
 49        seed_all(seed)
 50        d = dataset(seed)
 51        net = make_model("rnn_small", d["input_shape"], 3)
 52        device = "cuda" if torch.cuda.is_available() else "cpu"
 53        try:
 54            net = net.to(device)
 55            xtr, ytr = d["xtr"].to(device), d["ytr"].to(device)
 56            xte, yte = d["xte"].to(device), d["yte"].to(device)
 57            opt = torch.optim.Adam(net.parameters(), lr=cfg["lr"])
 58            for _ in range(EPOCHS):
 59                net.train()
 60                perm = torch.randperm(len(xtr), device=device)
 61                for j in range(0, len(xtr), BATCH):
 62                    ix = perm[j:j+BATCH]
 63                    raw = net(xtr[ix])
 64                    # Intrinsic S1 controller/state: retract the angular pair.
 65                    emb = raw[:, :2]
 66                    emb = emb / torch.clamp(torch.linalg.vector_norm(emb, dim=1, keepdim=True), min=1e-7)
 67                    pred = torch.cat((emb, raw[:, 2:3]), dim=1)
 68                    loss = ((pred - ytr[ix]) ** 2).mean()
 69                    opt.zero_grad(set_to_none=True)
 70                    loss.backward()
 71                    opt.step()
 72            net.eval()
 73            with torch.no_grad():
 74                raw = net(xte)
 75                emb = raw[:, :2] / torch.clamp(torch.linalg.vector_norm(raw[:, :2], dim=1, keepdim=True), min=1e-7)
 76                pred = torch.cat((emb, raw[:, 2:3]), dim=1)
 77                metric = float(((pred - yte) ** 2).mean())
 78                violation = float(torch.abs(torch.linalg.vector_norm(pred[:, :2], dim=1) - 1).max())
 79            if collect:
 80                return metric, net, d, violation
 81            return metric
 82        except RuntimeError:
 83            # Explicit CPU fallback, matching the bench's robust device policy.
 84            seed_all(seed)
 85            net = make_model("rnn_small", d["input_shape"], 3).to("cpu")
 86            xtr, ytr = d["xtr"], d["ytr"]
 87            opt = torch.optim.Adam(net.parameters(), lr=cfg["lr"])
 88            for _ in range(EPOCHS):
 89                perm = torch.randperm(len(xtr))
 90                for j in range(0, len(xtr), BATCH):
 91                    ix = perm[j:j+BATCH]
 92                    raw = net(xtr[ix]); emb = raw[:, :2] / torch.clamp(torch.linalg.vector_norm(raw[:, :2], dim=1, keepdim=True), min=1e-7)
 93                    pred = torch.cat((emb, raw[:, 2:3]), 1)
 94                    loss = ((pred-ytr[ix])**2).mean()
 95                    opt.zero_grad(set_to_none=True); loss.backward(); opt.step()
 96            with torch.no_grad():
 97                raw = net(d["xte"]); emb = raw[:, :2] / torch.clamp(torch.linalg.vector_norm(raw[:, :2], dim=1, keepdim=True), min=1e-7)
 98                pred = torch.cat((emb, raw[:, 2:3]), 1)
 99                metric = float(((pred-d["yte"])**2).mean())
100                violation = float(torch.abs(torch.linalg.vector_norm(pred[:, :2], dim=1)-1).max())
101            return (metric, net, d, violation) if collect else metric
102    return run
103
104
105def signature(cfg):
106    # Evaluate both trained systems on the same held-out examples.
107    s = 0
108    seed_all(s); d = dataset(s)
109    bnet, _, _ = train_model(make_model("rnn_small", d["input_shape"], 3), d,
110                             epochs=EPOCHS, lr=cfg["lr"], batch=BATCH, log=lambda *_: None)
111    bdev = next(bnet.parameters()).device
112    with torch.no_grad():
113        raw = bnet(d["xte"].to(bdev))
114        bviol = float(torch.abs(torch.linalg.vector_norm(raw[:, :2], dim=1)-1).max())
115    imetric, _, _, iviol = intrinsic_factory(cfg, collect=True)(s)
116    return {
117        "prediction": "intrinsic retraction keeps every predicted embedded angular state on S1; Euclidean output has nonzero norm error",
118        "predicted_baseline_violation_order": "nonzero",
119        "predicted_idea_violation": 0.0,
120        "observed_baseline_max_violation": bviol,
121        "observed_idea_max_violation": iviol,
122        "observed_idea_metric_seed0": imetric,
123        "confirmed": bool(bviol > 1e-6 and iviol < 1e-5)
124    }
125
126
127def main():
128    grid = [{"lr": x} for x in LRS]
129    base = sweep_baseline(baseline_factory, grid, seeds=(0, 1, 2, 3))
130    trials = [{"cfg": c, "result": evaluate(intrinsic_factory(c), SEEDS)} for c in grid]
131    best = min(trials, key=lambda z: z["result"]["mean"])
132    extra = {
133        "custom_track": {"name": "manifold_pendulum", "file": "manifold_dynamics_track.py", "domain": "dynamics_and_embedded_manifolds"},
134        "idea_config": best["cfg"],
135        "idea_sweep": trials,
136        "mechanism_signature": signature(best["cfg"])
137    }
138    rep = make_report("manifold_pendulum", "rnn_small", base, best["result"], extra)
139    with open("bench_report.json", "w") as f:
140        json.dump(rep, f, indent=2)
141    print(json.dumps(rep, indent=2))
142
143if __name__ == "__main__":
144    main()