Graded Levy-area recurrent state / stage2_bench.py

Mechanism confirmed, baseline not beaten

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