Observer-Corrected Robust Optimizer / bench_observer.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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()