import sys, json, random import numpy as np sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, make_report from bench.protocol import evaluate, sweep_baseline import torch from torch import nn DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu' class BurstGRU(nn.Module): def __init__(self, hidden=64, controlled=False, rc=1.5, delta=.25, gain=.15, qtol=.35): super().__init__() self.cell = nn.GRUCell(3, hidden) self.head = nn.Linear(hidden, 1) self.controlled = controlled self.rc, self.delta, self.gain, self.qtol = rc, delta, gain, qtol self.last_stats = {} def step(self, x, h): raw = self.cell(x, h) r = raw.norm(dim=1) u = raw / (r[:, None] + 1e-8) dv = raw - h a = (u * dv).sum(1) q = (dv - a[:, None] * u).norm(dim=1) gate = (r > self.rc) & (a > 0) & (q < self.qtol) if self.controlled: kappa = self.gain * torch.nn.functional.softplus((r-self.rc)/self.delta) raw = raw - gate[:, None] * kappa[:, None] * u self.last_stats = {'active': int(gate.sum().item()), 'max_r': float(r.detach().max().cpu()), 'mean_a': float(a.detach().mean().cpu()), 'mean_q': float(q.detach().mean().cpu())} return raw, gate, r, a, q def forward(self, x, collect=False): x = x.view(x.shape[0], -1, 3) h = torch.zeros(x.shape[0], self.cell.hidden_size, device=x.device) active = 0; maxr = 0.; apos = 0; near = 0; steps = 0 for k in range(x.shape[1]): h, gate, r, a, q = self.step(x[:, k, :], h) active += int(gate.sum().item()); maxr = max(maxr, float(r.max().detach().cpu())) apos += int((a > 0).sum().item()); near += int((q < self.qtol).sum().item()); steps += x.shape[0] if collect: return self.head(h), {'interventions': active, 'max_hidden_norm': maxr, 'positive_radial': apos/steps, 'near_equilibrium': near/steps} return self.head(h) def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def run(seed, cfg, controlled, collect=False): seed_all(seed) ds = get_dataset('dynamics', seed, n_train=400, n_test=400) try: dev = torch.device(DEVICE) model = BurstGRU(controlled=controlled, rc=cfg.get('rc',1.5), delta=cfg.get('delta',.25), gain=cfg.get('gain',.15), qtol=cfg.get('qtol',.35)).to(dev) opt = torch.optim.Adam(model.parameters(), lr=cfg['lr'], weight_decay=cfg.get('weight_decay',0.0)) lossf = nn.MSELoss() xtr,ytr,xte,yte = [ds[k].to(dev) for k in ('xtr','ytr','xte','yte')] model.train() for _ in range(cfg['epochs']): perm = torch.randperm(len(xtr), device=dev) for ix in perm.split(128): opt.zero_grad(set_to_none=True); pred=model(xtr[ix]); loss=lossf(pred,ytr[ix]); loss.backward(); opt.step() model.eval() with torch.no_grad(): pred, stats = model(xte, collect=True) mse=float(lossf(pred,yte).cpu()) return (mse, stats, model) if collect else mse except Exception: if DEVICE != 'cpu': old=globals()['DEVICE']; globals()['DEVICE']='cpu' try: return run(seed,cfg,controlled,collect) finally: globals()['DEVICE']=old raise def main(): # Union is shared: every idea lr is also baseline-tested. Method knobs are fixed a priori. grid=[{'lr':lr,'epochs':12,'weight_decay':wd} for lr in (0.001,0.003,0.009) for wd in (0.0,1e-4)] def baseline_fn(cfg): return lambda seed: run(seed,cfg,False) base=sweep_baseline(baseline_fn, grid, seeds=tuple(range(4))) cfgs=[base['best_cfg'], {'lr':0.001,'epochs':12,'weight_decay':base['best_cfg']['weight_decay']}, {'lr':0.009,'epochs':12,'weight_decay':base['best_cfg']['weight_decay']}] ideas=[] for cfg in cfgs: r=evaluate(lambda seed: run(seed,cfg,True), seeds=tuple(range(8))) ideas.append({'cfg':cfg,'result':r}) best=min(ideas,key=lambda z:z['result']['mean']) report=make_report('dynamics','rnn_small',base,best['result'],extra={'mechanism_signature': signature(base['best_cfg'],best['cfg'])}) report['idea_sweep']=ideas report['device']=DEVICE with open('bench_report.json','w') as f: json.dump(report,f,indent=2) print(json.dumps(report,indent=2)) def signature(base_cfg, idea_cfg): rows=[] for seed in range(8): _,bs,_=run(seed,base_cfg,False,True); _,ins,_=run(seed,idea_cfg,True,True) rows.append({'seed':seed,'baseline_max_hidden_norm':bs['max_hidden_norm'],'idea_max_hidden_norm':ins['max_hidden_norm'],'idea_interventions':ins['interventions'],'baseline_positive_radial':bs['positive_radial'],'idea_positive_radial':ins['positive_radial']}) b=np.mean([r['baseline_max_hidden_norm'] for r in rows]); i=np.mean([r['idea_max_hidden_norm'] for r in rows]); interventions=sum(r['idea_interventions'] for r in rows) return {'predicted_effect':'radial damping lowers burst norm when positive radial growth and near-equilibrium gate coincide','observed_mean_baseline_max_norm':float(b),'observed_mean_idea_max_norm':float(i),'mean_norm_reduction':float(b-i),'total_interventions':int(interventions),'confirmed':bool(interventions>0 and i