Compactified Burst Controller / compactified_bench.py
Mechanism confirmed, baseline not beaten
1import sys, json, random
2import numpy as np
3sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
4from bench import get_dataset, make_report
5from bench.protocol import evaluate, sweep_baseline
6import torch
7from torch import nn
8
9DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
10
11class BurstGRU(nn.Module):
12 def __init__(self, hidden=64, controlled=False, rc=1.5, delta=.25, gain=.15, qtol=.35):
13 super().__init__()
14 self.cell = nn.GRUCell(3, hidden)
15 self.head = nn.Linear(hidden, 1)
16 self.controlled = controlled
17 self.rc, self.delta, self.gain, self.qtol = rc, delta, gain, qtol
18 self.last_stats = {}
19
20 def step(self, x, h):
21 raw = self.cell(x, h)
22 r = raw.norm(dim=1)
23 u = raw / (r[:, None] + 1e-8)
24 dv = raw - h
25 a = (u * dv).sum(1)
26 q = (dv - a[:, None] * u).norm(dim=1)
27 gate = (r > self.rc) & (a > 0) & (q < self.qtol)
28 if self.controlled:
29 kappa = self.gain * torch.nn.functional.softplus((r-self.rc)/self.delta)
30 raw = raw - gate[:, None] * kappa[:, None] * u
31 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())}
32 return raw, gate, r, a, q
33
34 def forward(self, x, collect=False):
35 x = x.view(x.shape[0], -1, 3)
36 h = torch.zeros(x.shape[0], self.cell.hidden_size, device=x.device)
37 active = 0; maxr = 0.; apos = 0; near = 0; steps = 0
38 for k in range(x.shape[1]):
39 h, gate, r, a, q = self.step(x[:, k, :], h)
40 active += int(gate.sum().item()); maxr = max(maxr, float(r.max().detach().cpu()))
41 apos += int((a > 0).sum().item()); near += int((q < self.qtol).sum().item()); steps += x.shape[0]
42 if collect:
43 return self.head(h), {'interventions': active, 'max_hidden_norm': maxr, 'positive_radial': apos/steps, 'near_equilibrium': near/steps}
44 return self.head(h)
45
46def seed_all(seed):
47 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
48 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
49
50def run(seed, cfg, controlled, collect=False):
51 seed_all(seed)
52 ds = get_dataset('dynamics', seed, n_train=400, n_test=400)
53 try:
54 dev = torch.device(DEVICE)
55 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)
56 opt = torch.optim.Adam(model.parameters(), lr=cfg['lr'], weight_decay=cfg.get('weight_decay',0.0))
57 lossf = nn.MSELoss()
58 xtr,ytr,xte,yte = [ds[k].to(dev) for k in ('xtr','ytr','xte','yte')]
59 model.train()
60 for _ in range(cfg['epochs']):
61 perm = torch.randperm(len(xtr), device=dev)
62 for ix in perm.split(128):
63 opt.zero_grad(set_to_none=True); pred=model(xtr[ix]); loss=lossf(pred,ytr[ix]); loss.backward(); opt.step()
64 model.eval()
65 with torch.no_grad():
66 pred, stats = model(xte, collect=True)
67 mse=float(lossf(pred,yte).cpu())
68 return (mse, stats, model) if collect else mse
69 except Exception:
70 if DEVICE != 'cpu':
71 old=globals()['DEVICE']; globals()['DEVICE']='cpu'
72 try: return run(seed,cfg,controlled,collect)
73 finally: globals()['DEVICE']=old
74 raise
75
76def main():
77 # Union is shared: every idea lr is also baseline-tested. Method knobs are fixed a priori.
78 grid=[{'lr':lr,'epochs':12,'weight_decay':wd} for lr in (0.001,0.003,0.009) for wd in (0.0,1e-4)]
79 def baseline_fn(cfg): return lambda seed: run(seed,cfg,False)
80 base=sweep_baseline(baseline_fn, grid, seeds=tuple(range(4)))
81 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']}]
82 ideas=[]
83 for cfg in cfgs:
84 r=evaluate(lambda seed: run(seed,cfg,True), seeds=tuple(range(8)))
85 ideas.append({'cfg':cfg,'result':r})
86 best=min(ideas,key=lambda z:z['result']['mean'])
87 report=make_report('dynamics','rnn_small',base,best['result'],extra={'mechanism_signature': signature(base['best_cfg'],best['cfg'])})
88 report['idea_sweep']=ideas
89 report['device']=DEVICE
90 with open('bench_report.json','w') as f: json.dump(report,f,indent=2)
91 print(json.dumps(report,indent=2))
92
93def signature(base_cfg, idea_cfg):
94 rows=[]
95 for seed in range(8):
96 _,bs,_=run(seed,base_cfg,False,True); _,ins,_=run(seed,idea_cfg,True,True)
97 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']})
98 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)
99 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<b)}
100
101if __name__=='__main__': main()