Observer-Corrected Robust Optimizer / bench_observer.py
Mechanism confirmed, baseline not beaten
1"""Stage-2 benchmark for observer-corrected SGD on the matched dynamics track."""
2import json, random
3from pathlib import Path
4import numpy as np
5import torch
6import torch.nn as nn
7
8import sys
9sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
10from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
11
12SEEDS = tuple(range(8))
13# Union of all learning rates is shared by baseline and idea.
14LRS = [1e-3, 3e-3, 1e-2]
15MOMENTA = [0.0, 0.9]
16EPOCHS = 18
17BATCH = 64
18
19
20def seed_all(seed):
21 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
22 if torch.cuda.is_available():
23 try: torch.cuda.manual_seed_all(seed)
24 except Exception: pass
25
26
27def device_ladder():
28 if torch.cuda.is_available(): return ["cuda", "cpu"]
29 return ["cpu"]
30
31
32def loss_fn(ds):
33 return nn.CrossEntropyLoss() if ds["task"] == "classification" else nn.MSELoss()
34
35
36def baseline_train(seed, cfg, collect=False):
37 seed_all(seed); ds = get_dataset("dynamics", seed, n_train=400, n_test=400)
38 net = make_model("rnn_small", ds["input_shape"], ds["out_dim"])
39 lf = loss_fn(ds); lr, mu = cfg["lr"], cfg["momentum"]
40 last_sig = {}
41 for dev in device_ladder():
42 try:
43 net = net.to(dev); x, y = ds["xtr"].to(dev), ds["ytr"].to(dev)
44 opt = torch.optim.SGD(net.parameters(), lr=lr, momentum=mu)
45 for _ in range(EPOCHS):
46 net.train(); perm = torch.randperm(len(x), device=dev)
47 for j in range(0, len(x), BATCH):
48 ix = perm[j:j+BATCH]; z = lf(net(x[ix]), y[ix])
49 opt.zero_grad(); z.backward(); opt.step()
50 net.eval()
51 with torch.no_grad(): metric = float(lf(net(ds["xte"].to(dev)), ds["yte"].to(dev)))
52 return metric, last_sig
53 except RuntimeError:
54 net = make_model("rnn_small", ds["input_shape"], ds["out_dim"])
55 return float("nan"), last_sig
56
57
58def observer_train(seed, cfg, collect=False):
59 """SGD momentum plus EMA of one-step gradient-transition residual.
60
61 g_ema predicts the slowly varying nominal gradient. The residual observer
62 tracks r=g-g_ema and subtracts its EMA from the momentum direction.
63 The correction is norm-clipped, which is the robust bounded-gain safeguard.
64 """
65 seed_all(seed); ds = get_dataset("dynamics", seed, n_train=400, n_test=400)
66 net = make_model("rnn_small", ds["input_shape"], ds["out_dim"])
67 lf = loss_fn(ds); lr, mu = cfg["lr"], cfg["momentum"]
68 alpha, clip = cfg["alpha"], cfg["clip"]
69 dhat = [torch.zeros_like(p) for p in net.parameters()]
70 gprev = [torch.zeros_like(p) for p in net.parameters()]
71 raw_sq = corr_sq = applied_sq = 0.0; count = 0
72 for dev in device_ladder():
73 try:
74 net = net.to(dev); x, y = ds["xtr"].to(dev), ds["ytr"].to(dev)
75 dhat = [q.to(dev) for q in dhat]; gprev = [q.to(dev) for q in gprev]
76 velocity = [torch.zeros_like(p) for p in net.parameters()]
77 for _ in range(EPOCHS):
78 net.train(); perm = torch.randperm(len(x), device=dev)
79 for j in range(0, len(x), BATCH):
80 ix = perm[j:j+BATCH]
81 z = lf(net(x[ix]), y[ix]); grads = torch.autograd.grad(z, tuple(net.parameters()))
82 with torch.no_grad():
83 for p,v,gh,old in zip(net.parameters(), velocity, grads, gprev):
84 # transition residual: observed gradient minus prior prediction
85 r = gh - old
86 newd = (1-alpha) * dhat[len([q for q in []])] if False else None
87 # indexed loop avoids hidden optimizer state
88 for k,(p,v,gh,old) in enumerate(zip(net.parameters(), velocity, grads, gprev)):
89 r = gh - old
90 dhat[k].mul_(1-alpha).add_(r, alpha=alpha)
91 n = torch.linalg.vector_norm(dhat[k])
92 if n > clip: dhat[k].mul_(clip/(n+1e-12))
93 v.mul_(mu).add_(gh)
94 applied = v - dhat[k]
95 p.add_(applied, alpha=-lr)
96 if collect:
97 raw_sq += float(torch.sum(r*r)); corr_sq += float(torch.sum(dhat[k]*dhat[k])); applied_sq += float(torch.sum(applied*applied)); count += r.numel()
98 old.copy_(gh)
99 net.eval()
100 with torch.no_grad(): metric = float(lf(net(ds["xte"].to(dev)), ds["yte"].to(dev)))
101 sig = {"raw_transition_rms": float(np.sqrt(raw_sq/max(count,1))), "observer_rms": float(np.sqrt(corr_sq/max(count,1))), "applied_direction_rms": float(np.sqrt(applied_sq/max(count,1)))}
102 return metric, sig
103 except RuntimeError:
104 net = make_model("rnn_small", ds["input_shape"], ds["out_dim"])
105 return float("nan"), {}
106
107
108def main():
109 # Baseline sweep includes every lr used by idea and the central SGD knob momentum.
110 grid = [{"lr": lr, "momentum": mu} for lr in LRS for mu in MOMENTA]
111 base = sweep_baseline(lambda c: lambda s: baseline_train(s,c)[0], grid, seeds=(0,1,2,3))
112 # Three observer settings, same lr/momentum union and equal 3-config idea budget.
113 best = base["best_cfg"]
114 idea_grid = [{"lr": best["lr"], "momentum": best["momentum"], "alpha": a, "clip": 0.05} for a in (0.02,0.1,0.3)]
115 # Also ensure nearby learning rates are represented, while baseline already evaluated them.
116 idea_grid[0]["lr"] = LRS[max(0,LRS.index(best["lr"])-1)]
117 idea_grid[2]["lr"] = LRS[min(len(LRS)-1,LRS.index(best["lr"])+1)]
118 results=[]
119 for c in idea_grid:
120 r=evaluate(lambda s: observer_train(s,c)[0], seeds=SEEDS)
121 results.append((r,c))
122 idea,cfg=max(results, key=lambda q: -q[0]["mean"] if np.isfinite(q[0]["mean"]) else -1e99)
123 # Correct selection is minimum metric.
124 idea,cfg=min(results, key=lambda q: q[0]["mean"])
125 sigs=[observer_train(s,cfg,True)[1] for s in SEEDS]
126 sig={k: float(np.mean([x[k] for x in sigs])) for k in sigs[0]}
127 sig.update({"predicted_residual_reduction": float(1-sig["observer_rms"]/max(sig["raw_transition_rms"],1e-12)), "confirmed": bool(sig["observer_rms"] < sig["raw_transition_rms"])})
128 rep=make_report("dynamics","rnn_small",base,idea,{"prediction":"EMA observer reduces slowly varying transition residual; measured on trained models","values":sig,"confirmed":sig["confirmed"]})
129 rep["idea_sweep"]=[{"cfg":c,"mean":r["mean"]} for r,c in results]
130 rep["matched_structure_justification"]="Dynamics is the control/stability track; both systems train identical rnn_small models on identical paired pendulum datasets."
131 Path("bench_report.json").write_text(json.dumps(rep,indent=2))
132 print(json.dumps(rep,indent=2))
133
134if __name__ == "__main__": main()