import os, sys, json, random, math import numpy as np import torch import torch.nn as nn sys.path.insert(0, "/home/maxwelhelp/all/math2nn") from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report, permutation_pvalue, count_params TRACK = "sequence" EPOCHS = 15 NTRAIN, NTEST = 400, 400 SEEDS = tuple(range(8)) SWEEP_SEEDS = tuple(range(4)) class AdditiveState(nn.Module): def __init__(self, input_shape, out_dim, k=16): super().__init__() self.k = k self.proj = nn.Linear(1, k) self.head = nn.Linear(k + k, out_dim) def forward(self, x): v = torch.tanh(self.proj(x.unsqueeze(-1))) u = v.sum(1) n = float(x.shape[1]) feat = torch.cat((u / math.sqrt(n), (v*v).sum(1) / n), 1) return self.head(feat) class LevyState(nn.Module): def __init__(self, input_shape, out_dim, k=16): super().__init__() self.k = k self.proj = nn.Linear(1, k) self.iu = torch.triu_indices(k, k, 1) self.head = nn.Linear(k + k*(k-1)//2 + k, out_dim) def forward(self, x, return_state=False): v = torch.tanh(self.proj(x.unsqueeze(-1))) u = torch.zeros(x.shape[0], self.k, device=x.device, dtype=x.dtype) A = torch.zeros(x.shape[0], self.k, self.k, device=x.device, dtype=x.dtype) q = torch.zeros_like(u) for t in range(x.shape[1]): z = v[:, t] A = A + 0.5 * (u[:,:,None] * z[:,None,:] - z[:,:,None] * u[:,None,:]) u = u + z q = q + z*z n = float(x.shape[1]) feat = torch.cat((u / math.sqrt(n), A[:, self.iu[0], self.iu[1]] / n, q / n), 1) out = self.head(feat) if return_state: return out, (u, A, q, feat) return out def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) def train_one(kind, cfg, seed, return_model=False): seed_all(seed) ds = get_dataset(TRACK, seed, n_train=NTRAIN, n_test=NTEST) cls = AdditiveState if kind == "baseline" else LevyState model = cls(ds["input_shape"], ds["out_dim"], k=cfg.get("k", 16)) net, metric, hist = train_model(model, ds, epochs=EPOCHS, lr=cfg["lr"], batch=128, log=lambda *_: None) if net is None: raise RuntimeError("training failed") if return_model: return float(metric), net, ds return float(metric) def make_fn(kind): return lambda cfg: (lambda seed: train_one(kind, cfg, seed)) def signature(seed, cfg): metric, model, ds = train_one("idea", cfg, seed, True) model.eval(); dev = next(model.parameters()).device; x = ds["xte"][:64].to(dev) with torch.no_grad(): observed = (model(x)[...,0] - model(torch.flip(x, dims=[1]))[...,0]).cpu().numpy() _, (_, A, _, _) = model(x, True) _, (_, Ar, _, _) = model(torch.flip(x, dims=[1]), True) # Apply the trained head to the area-only feature difference, retaining # the learned projection and head from this trained benchmark model. n=float(x.shape[1]); du=torch.zeros_like(A); dq=torch.zeros_like(A[:,0,0:1]) z=torch.zeros(x.shape[0], model.k, device=x.device) # Difference of complete states isolates the area contribution because # u and q are reversal invariant; compute it directly through features. fa=torch.zeros(x.shape[0], model.k*(model.k-1)//2, device=x.device) diffA=(A-Ar) / n fa=diffA[:, model.iu[0], model.iu[1]] w=model.head.weight[0, model.k:model.k+fa.shape[1]] predicted=(fa*w).sum(1).cpu().numpy() corr=float(np.corrcoef(observed, predicted)[0,1]) if np.std(predicted)>1e-12 else 0.0 ratio=float(np.linalg.norm(observed)/(np.linalg.norm(predicted)+1e-12)) return {"n": int(len(observed)), "observed_output_reversal_rms": float(np.sqrt(np.mean(observed**2))), "predicted_area_head_rms": float(np.sqrt(np.mean(predicted**2))), "observed_to_predicted_rms_ratio": ratio, "correlation": corr, "confirmed": bool(corr > .95 and .8 < ratio < 1.2), "model_test_mse": metric} def main(): # Union parity: every LR used by either side is swept on both sides. grid=[{"lr":1e-3},{"lr":3e-3},{"lr":1e-2}] base=sweep_baseline(make_fn("baseline"), grid, seeds=SWEEP_SEEDS) idea_sweep=sweep_baseline(make_fn("idea"), grid, seeds=SWEEP_SEEDS) best=idea_sweep["best_cfg"] idea_full=evaluate(make_fn("idea")(best), seeds=SEEDS) sig=signature(0, best) report=make_report(TRACK, "custom_recurrent_state", base, idea_full, {"prediction": "reversal output difference is explained by trained area-head contribution", **sig, "idea_sweep": idea_sweep}) report["protocol_notes"]={"architecture": "matched projection/state/readout; only area channels replace additive state", "epochs":EPOCHS,"n_train":NTRAIN,"n_test":NTEST,"paired_seeds":list(SEEDS), "baseline_and_idea_lr_union":grid,"baseline_selection_seeds":list(SWEEP_SEEDS), "parameters_baseline":count_params(AdditiveState((32,),1)),"parameters_idea":count_params(LevyState((32,),1))} with open("bench_report.json","w") as f: json.dump(report,f,indent=2) print(json.dumps(report,indent=2)) if __name__ == "__main__": main()