import os, sys, json, math, random from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, "/home/maxwelhelp/all/math2nn") from bench import get_dataset, sweep_baseline, evaluate, make_report OUT = Path(__file__).with_name("bench_report.json") SEEDS = tuple(range(8)) SWEEP_SEEDS = (0, 1, 2, 3) EPOCHS = 12 NTRAIN, NTEST = 1200, 400 BATCH = 128 # The union is deliberately shared by both sides. LRS = (1e-3, 3e-3, 1e-2) IDEA_GRID = ({"lr": 1e-3, "k": 0.5, "delta": 0.02}, {"lr": 3e-3, "k": 0.5, "delta": 0.05}, {"lr": 1e-2, "k": 0.5, "delta": 0.10}) BASE_GRID = [{"lr": x, "weight_decay": wd} for x in LRS for wd in (0.0, 1e-4)] class MatchedGRU(nn.Module): """GRUCell(3,64)+linear, optionally followed by phase control.""" def __init__(self, out_dim=1, mode="baseline", k=0.5, delta=0.05): super().__init__() self.cell = nn.GRUCell(3, 64) self.head = nn.Linear(64, out_dim) self.mode, self.k, self.delta = mode, k, delta self.last_stats = {} @staticmethod def control(h, k): p = h.view(h.shape[0], 32, 2) z = torch.complex(p[..., 0], p[..., 1]) z = z / (torch.abs(z) + 1e-6) r = z.mean(1, keepdim=True) return (2*k/32) * torch.imag(z * torch.conj(r)) @staticmethod def rotate(h, u, dt=1.0): p = h.view(h.shape[0], 32, 2) a = dt*u c, s = torch.cos(a), torch.sin(a) x, y = p[..., 0], p[..., 1] return torch.stack((x*c-y*s, x*s+y*c), -1).reshape_as(h) def forward(self, x, collect=False): seq = x.view(x.shape[0], -1, 3) h = torch.zeros(x.shape[0], 64, device=x.device, dtype=x.dtype) held = torch.zeros(x.shape[0], 32, device=x.device, dtype=x.dtype) events = 0 vs, exacts, holds, dv_exact, dv_held = [], [], [], [], [] for t in range(seq.shape[1]): h = self.cell(seq[:, t], h) if self.mode != "baseline": ue = self.control(h, self.k) # Event logic is a controller implementation detail; no STE. trigger = torch.linalg.vector_norm(ue - held, dim=1) >= self.delta held = torch.where(trigger[:, None], ue, held) events += int(trigger.sum().item()) if collect: ph = h.view(h.shape[0], 32, 2) z0 = torch.complex(ph[...,0], ph[...,1]); z0 = z0/(torch.abs(z0)+1e-6) v0 = torch.abs(z0.mean(1))**2 he = self.rotate(h, ue, 1.0); hh = self.rotate(h, held, 1.0) def vv(q): qp=q.view(q.shape[0],32,2); zz=torch.complex(qp[...,0],qp[...,1]); zz=zz/(torch.abs(zz)+1e-6); return torch.abs(zz.mean(1))**2 dv_exact.append((vv(he)-v0).detach()); dv_held.append((vv(hh)-v0).detach()) z = torch.complex(h.view(h.shape[0],32,2)[...,0], h.view(h.shape[0],32,2)[...,1]) z = z/(torch.abs(z)+1e-6) vs.append(torch.abs(z.mean(1))**2) exacts.append(ue.detach()); holds.append(held.detach()) h = self.rotate(h, held, 1.0) self.last_stats = {"events": events, "steps": x.shape[0]*seq.shape[1], "v": vs, "u": exacts, "held": holds, "dv_exact": dv_exact, "dv_held": dv_held} return self.head(h) def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) def run_one(seed, cfg, mode, collect=False): seed_all(seed) d = get_dataset("dynamics", seed, n_train=NTRAIN, n_test=NTEST) model = MatchedGRU(mode=mode, k=cfg.get("k",0.0), delta=cfg.get("delta",1e9)) dev = torch.device("cuda" if torch.cuda.is_available() else "cpu") try: model.to(dev); xtr,ytr=d["xtr"].to(dev),d["ytr"].to(dev) xte,yte=d["xte"].to(dev),d["yte"].to(dev) opt=torch.optim.Adam(model.parameters(), lr=cfg["lr"], weight_decay=cfg.get("weight_decay",0.0)) for ep in range(EPOCHS): model.train(); perm=torch.randperm(len(xtr),device=dev) for i in range(0,len(xtr),BATCH): ix=perm[i:i+BATCH]; loss=((model(xtr[ix])-ytr[ix])**2).mean() opt.zero_grad(); loss.backward(); opt.step() model.eval() with torch.no_grad(): pred=model(xte, collect=collect); metric=((pred-yte)**2).mean().item() stats=model.last_stats result={"metric":metric,"events_per_sequence":stats["events"]/NTEST, "model":model,"stats":stats,"device":str(dev)} return result except RuntimeError: if dev.type == "cuda": torch.cuda.empty_cache(); os.environ["CUDA_VISIBLE_DEVICES"]="" return run_one(seed,cfg,mode,collect) raise def metric_runner(mode, cfg): return lambda seed: run_one(seed,cfg,mode)["metric"] def mechanism_signature(cfg): # Both values are measured from hidden states produced by trained models. pred, exact_obs, held_obs, errs, counts = [], [], [], [], [] for seed in SEEDS: r = run_one(seed, cfg, "event", collect=True); st = r["stats"] for u, de, dh in zip(st["u"], st["dv_exact"], st["dv_held"]): # First-order prediction for dt=1 under the exact tangent field. pred.append((-cfg["k"] * ((u/cfg["k"])**2).sum(1)).mean().item()) exact_obs.append(de.mean().item()); held_obs.append(dh.mean().item()) for u, h in zip(st["u"], st["held"]): errs.append(torch.linalg.vector_norm(u-h, dim=1).max().item()) counts.append(st["events"]) pm, em, hm = map(float, (np.mean(pred), np.mean(exact_obs), np.mean(held_obs))) ratio = em/pm if pm else float("nan") return {"quantity":"one-step V change on trained hidden states", "predicted_exact_delta_V":pm, "observed_exact_delta_V":em, "observed_held_delta_V":hm, "exact_prediction_ratio":ratio, "max_control_hold_error":float(max(errs)), "mean_events_per_sequence":float(np.mean(counts)/NTEST), "confirmed":bool(np.isfinite(ratio) and abs(ratio-1)<0.10)} def main(): # cheap core math check first: finite difference of V under the exact law rng=np.random.default_rng(123); th=rng.normal(0,.25,32); z=np.exp(1j*th); r=z.mean(); u=2/32*np.imag(z*np.conj(r)); eps=1e-6 v0=abs(r)**2; v1=abs(np.exp(1j*(th+eps*u)).mean())**2 math_check={"observed_vdot":float((v1-v0)/eps),"predicted_vdot":float(-np.sum(u*u)), "ratio":float(((v1-v0)/eps)/(-np.sum(u*u)))} base=sweep_baseline(lambda c: metric_runner("baseline",c), BASE_GRID, seeds=SWEEP_SEEDS) best_lr=base["best_cfg"]["lr"] # Idea is evaluated at best baseline lr and two nearby union-grid settings. idea_cfgs=[dict(c) for c in IDEA_GRID] idea_cfgs[1]["lr"]=best_lr idea_sweep=[] for c in idea_cfgs: rr=evaluate(metric_runner("event",c), seeds=SWEEP_SEEDS) idea_sweep.append({"cfg":c,"mean":rr["mean"]}) best_idea=min(idea_sweep,key=lambda q:q["mean"])["cfg"] idea=evaluate(metric_runner("event",best_idea), seeds=SEEDS) sig=mechanism_signature(best_idea) report=make_report("dynamics","rnn_small",base,idea, {"mechanism_signature":sig,"math_check":math_check, "idea_sweep":idea_sweep,"protocol":{"epochs":EPOCHS,"n_train":NTRAIN,"n_test":NTEST,"paired_seeds":list(SEEDS)}}) OUT.write_text(json.dumps(report,indent=2,allow_nan=False)) print(json.dumps(report,indent=2)) if __name__=="__main__": main()