Graded Levy-area recurrent state / stage2_bench.py
Mechanism confirmed, baseline not beaten
1import os, sys, json, random, math
2import numpy as np
3import torch
4import torch.nn as nn
5
6sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
7from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report, permutation_pvalue, count_params
8
9TRACK = "sequence"
10EPOCHS = 15
11NTRAIN, NTEST = 400, 400
12SEEDS = tuple(range(8))
13SWEEP_SEEDS = tuple(range(4))
14
15class AdditiveState(nn.Module):
16 def __init__(self, input_shape, out_dim, k=16):
17 super().__init__()
18 self.k = k
19 self.proj = nn.Linear(1, k)
20 self.head = nn.Linear(k + k, out_dim)
21 def forward(self, x):
22 v = torch.tanh(self.proj(x.unsqueeze(-1)))
23 u = v.sum(1)
24 n = float(x.shape[1])
25 feat = torch.cat((u / math.sqrt(n), (v*v).sum(1) / n), 1)
26 return self.head(feat)
27
28class LevyState(nn.Module):
29 def __init__(self, input_shape, out_dim, k=16):
30 super().__init__()
31 self.k = k
32 self.proj = nn.Linear(1, k)
33 self.iu = torch.triu_indices(k, k, 1)
34 self.head = nn.Linear(k + k*(k-1)//2 + k, out_dim)
35 def forward(self, x, return_state=False):
36 v = torch.tanh(self.proj(x.unsqueeze(-1)))
37 u = torch.zeros(x.shape[0], self.k, device=x.device, dtype=x.dtype)
38 A = torch.zeros(x.shape[0], self.k, self.k, device=x.device, dtype=x.dtype)
39 q = torch.zeros_like(u)
40 for t in range(x.shape[1]):
41 z = v[:, t]
42 A = A + 0.5 * (u[:,:,None] * z[:,None,:] - z[:,:,None] * u[:,None,:])
43 u = u + z
44 q = q + z*z
45 n = float(x.shape[1])
46 feat = torch.cat((u / math.sqrt(n), A[:, self.iu[0], self.iu[1]] / n, q / n), 1)
47 out = self.head(feat)
48 if return_state: return out, (u, A, q, feat)
49 return out
50
51def seed_all(seed):
52 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
53
54def train_one(kind, cfg, seed, return_model=False):
55 seed_all(seed)
56 ds = get_dataset(TRACK, seed, n_train=NTRAIN, n_test=NTEST)
57 cls = AdditiveState if kind == "baseline" else LevyState
58 model = cls(ds["input_shape"], ds["out_dim"], k=cfg.get("k", 16))
59 net, metric, hist = train_model(model, ds, epochs=EPOCHS, lr=cfg["lr"], batch=128, log=lambda *_: None)
60 if net is None: raise RuntimeError("training failed")
61 if return_model: return float(metric), net, ds
62 return float(metric)
63
64def make_fn(kind):
65 return lambda cfg: (lambda seed: train_one(kind, cfg, seed))
66
67def signature(seed, cfg):
68 metric, model, ds = train_one("idea", cfg, seed, True)
69 model.eval(); dev = next(model.parameters()).device; x = ds["xte"][:64].to(dev)
70 with torch.no_grad():
71 observed = (model(x)[...,0] - model(torch.flip(x, dims=[1]))[...,0]).cpu().numpy()
72 _, (_, A, _, _) = model(x, True)
73 _, (_, Ar, _, _) = model(torch.flip(x, dims=[1]), True)
74 # Apply the trained head to the area-only feature difference, retaining
75 # the learned projection and head from this trained benchmark model.
76 n=float(x.shape[1]); du=torch.zeros_like(A); dq=torch.zeros_like(A[:,0,0:1])
77 z=torch.zeros(x.shape[0], model.k, device=x.device)
78 # Difference of complete states isolates the area contribution because
79 # u and q are reversal invariant; compute it directly through features.
80 fa=torch.zeros(x.shape[0], model.k*(model.k-1)//2, device=x.device)
81 diffA=(A-Ar) / n
82 fa=diffA[:, model.iu[0], model.iu[1]]
83 w=model.head.weight[0, model.k:model.k+fa.shape[1]]
84 predicted=(fa*w).sum(1).cpu().numpy()
85 corr=float(np.corrcoef(observed, predicted)[0,1]) if np.std(predicted)>1e-12 else 0.0
86 ratio=float(np.linalg.norm(observed)/(np.linalg.norm(predicted)+1e-12))
87 return {"n": int(len(observed)), "observed_output_reversal_rms": float(np.sqrt(np.mean(observed**2))),
88 "predicted_area_head_rms": float(np.sqrt(np.mean(predicted**2))),
89 "observed_to_predicted_rms_ratio": ratio, "correlation": corr,
90 "confirmed": bool(corr > .95 and .8 < ratio < 1.2), "model_test_mse": metric}
91
92def main():
93 # Union parity: every LR used by either side is swept on both sides.
94 grid=[{"lr":1e-3},{"lr":3e-3},{"lr":1e-2}]
95 base=sweep_baseline(make_fn("baseline"), grid, seeds=SWEEP_SEEDS)
96 idea_sweep=sweep_baseline(make_fn("idea"), grid, seeds=SWEEP_SEEDS)
97 best=idea_sweep["best_cfg"]
98 idea_full=evaluate(make_fn("idea")(best), seeds=SEEDS)
99 sig=signature(0, best)
100 report=make_report(TRACK, "custom_recurrent_state", base, idea_full,
101 {"prediction": "reversal output difference is explained by trained area-head contribution",
102 **sig, "idea_sweep": idea_sweep})
103 report["protocol_notes"]={"architecture": "matched projection/state/readout; only area channels replace additive state",
104 "epochs":EPOCHS,"n_train":NTRAIN,"n_test":NTEST,"paired_seeds":list(SEEDS),
105 "baseline_and_idea_lr_union":grid,"baseline_selection_seeds":list(SWEEP_SEEDS),
106 "parameters_baseline":count_params(AdditiveState((32,),1)),"parameters_idea":count_params(LevyState((32,),1))}
107 with open("bench_report.json","w") as f: json.dump(report,f,indent=2)
108 print(json.dumps(report,indent=2))
109
110if __name__ == "__main__": main()