Intrinsic Schrödinger Bridge Diffusion / stage2_bench.py
Mechanism confirmed, baseline not beaten
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()