Gauge-Patched Local Experts / bench_gauge_dynamics.py

✓✓ Beats tuned baseline

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