Gauge-Patched Local Experts / bench_gauge_dynamics.py
Beats tuned baseline
1import json, sys, math
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6
7sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
8from bench import get_dataset, evaluate, sweep_baseline, make_report
9
10SEED0 = 1270
11N_TRAIN, N_TEST = 1200, 400
12EPOCHS, BATCH = 18, 128
13LR_GRID = [1e-3, 3e-3, 1e-2]
14LAMBDA_PATCH = 0.01
15DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
16
17class GaugeRNN(nn.Module):
18 """rnn_small-compatible GRU with two local hidden patches and SO(2)^32 gauges."""
19 def __init__(self, hidden=64):
20 super().__init__()
21 self.rnn = nn.GRU(3, hidden, batch_first=True)
22 self.head = nn.Linear(hidden, 1)
23 # One angle per 2D coordinate plane; exp(skew(angle)) is exactly orthogonal.
24 self.angles = nn.Parameter(torch.zeros(hidden // 2))
25
26 def forward(self, x, return_aux=False):
27 seq = x.view(x.shape[0], -1, 3)
28 _, hfull = self.rnn(seq)
29 pred = self.head(hfull[-1])
30 if not return_aux:
31 return pred
32 mid = seq.shape[1] // 2
33 _, h0 = self.rnn(seq[:, :mid])
34 _, h1 = self.rnn(seq[:, mid:])
35 a = self.angles
36 c, s = torch.cos(a), torch.sin(a)
37 z0, z1 = h0[-1].view(h0[-1].shape[0], self.angles.numel(), 2), h1[-1].view(h1[-1].shape[0], self.angles.numel(), 2)
38 # R(a) @ [x,y] = [c*x-s*y, s*x+c*y]
39 zg = torch.stack((c[None] * z1[..., 0] - s[None] * z1[..., 1],
40 s[None] * z1[..., 0] + c[None] * z1[..., 1]), dim=-1)
41 aligned = zg.reshape_as(h1[-1])
42 return pred, h0[-1], h1[-1], aligned
43
44def patch_loss(model, x):
45 _, h0, h1, aligned = model(x, return_aux=True)
46 return ((h0 - aligned) ** 2).mean()
47
48def train_one(seed, lr, lam):
49 global DEVICE
50 torch.manual_seed(seed); np.random.seed(seed)
51 try:
52 ds = get_dataset("dynamics", seed, n_train=N_TRAIN, n_test=N_TEST)
53 model = GaugeRNN().to(DEVICE)
54 opt = torch.optim.Adam(model.parameters(), lr=lr)
55 xtr, ytr = ds["xtr"].to(DEVICE), ds["ytr"].to(DEVICE)
56 model.train()
57 gen = torch.Generator(device="cpu").manual_seed(seed + 91)
58 for ep in range(EPOCHS):
59 order = torch.randperm(len(xtr), generator=gen)
60 for ix in order.split(BATCH):
61 xb, yb = xtr[ix], ytr[ix]
62 pred = model(xb)
63 task = ((pred - yb) ** 2).mean()
64 loss = task if lam == 0.0 else task + lam * patch_loss(model, xb)
65 opt.zero_grad(set_to_none=True); loss.backward(); opt.step()
66 model.eval()
67 with torch.no_grad():
68 pred, h0, h1, aligned = model(ds["xte"].to(DEVICE), return_aux=True)
69 metric = ((pred - ds["yte"].to(DEVICE)) ** 2).mean().item()
70 disagreement = ((h0 - aligned) ** 2).mean().item()
71 norms = torch.linalg.vector_norm(aligned, dim=1) - torch.linalg.vector_norm(h1, dim=1)
72 norm_err = norms.abs().max().item()
73 angle_mag = model.angles.abs().mean().item()
74 return metric, {"disagreement": disagreement, "norm_error": norm_err, "angle_mean_abs": angle_mag}
75 except RuntimeError as exc:
76 # CUDA failures get one deterministic CPU retry; programming errors propagate.
77 msg = str(exc).lower()
78 if DEVICE != "cuda" or not any(k in msg for k in ("cuda", "cudnn", "out of memory")):
79 raise
80 old = DEVICE; DEVICE = "cpu"
81 try: return train_one(seed, lr, lam)
82 finally: DEVICE = old
83
84def metric_fn(lam, lr):
85 def f(seed): return train_one(seed, lr, lam)[0]
86 return f
87
88def main():
89 # Cheap numerical verification before training: block rotations preserve norms.
90 torch.manual_seed(SEED0)
91 a = torch.randn(17, 32); theta = torch.linspace(-2, 2, 16)
92 c, s = torch.cos(theta), torch.sin(theta)
93 b = torch.stack((c[None] * a[:, 0::2] - s[None] * a[:, 1::2],
94 s[None] * a[:, 0::2] + c[None] * a[:, 1::2]), -1).reshape_as(a)
95 math_check = {"max_norm_error": float((a.norm(dim=1)-b.norm(dim=1)).abs().max()),
96 "predicted": "orthogonal gauge preserves hidden norm"}
97
98 # Baseline sweep uses every learning rate also tried by the idea.
99 base = sweep_baseline(lambda cfg: metric_fn(0.0, cfg["lr"]),
100 [{"lr": x} for x in LR_GRID])
101 idea_runs = []
102 for lr in LR_GRID:
103 r = evaluate(metric_fn(LAMBDA_PATCH, lr))
104 idea_runs.append({"cfg": {"lr": lr, "lambda_patch": LAMBDA_PATCH}, "result": r})
105 best = min(idea_runs, key=lambda z: z["result"]["mean"])
106 idea = best["result"]
107
108 # Re-test trained models on all paired seeds for mechanism signature.
109 base_beh, idea_beh = [], []
110 for s in range(8):
111 bm, bx = train_one(s, base["best_cfg"]["lr"], 0.0)
112 im, ix = train_one(s, best["cfg"]["lr"], LAMBDA_PATCH)
113 base_beh.append({"metric": bm, **bx}); idea_beh.append({"metric": im, **ix})
114 bdisc = float(np.mean([z["disagreement"] for z in base_beh]))
115 idisc = float(np.mean([z["disagreement"] for z in idea_beh]))
116 inorm = float(max(z["norm_error"] for z in idea_beh))
117 signature = {"prediction": "orthogonal transitions preserve hidden norms and patch loss reduces disagreement",
118 "predicted_norm_error": 0.0, "observed_max_norm_error": inorm,
119 "predicted_disagreement_change": "negative", "observed_baseline_disagreement": bdisc,
120 "observed_idea_disagreement": idisc,
121 "confirmed": bool(inorm < 1e-5 and idisc < bdisc),
122 "math_sanity": math_check, "trained_model_behavior": {"baseline": base_beh, "idea": idea_beh}}
123 rep = make_report("dynamics", "rnn_small", base, idea,
124 {**signature, "idea_sweep": idea_runs,
125 "track_rationale": "Dynamics is the built-in stability/control track; local temporal patches are neighboring state windows."})
126 Path("bench_report.json").write_text(json.dumps(rep, indent=2))
127 print(json.dumps(rep, indent=2))
128
129if __name__ == "__main__": main()