Gumbel escape-time controller / bench_gumbel_controller.py
Failed on benchmark
1import json, math, random
2import numpy as np
3import torch
4import torch.nn as nn
5from scipy.stats import gumbel_r, kstest
6import sys
7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
8from bench import get_dataset, make_model, sweep_baseline, evaluate, make_report
9
10SEEDS = tuple(range(8))
11NTRAIN, NTEST = 400, 400
12EPOCHS, BATCH = 12, 64
13
14
15def seed_all(seed):
16 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
17 if torch.cuda.is_available():
18 try: torch.cuda.manual_seed_all(seed)
19 except Exception: pass
20
21
22def device():
23 return 'cuda' if torch.cuda.is_available() else 'cpu'
24
25
26def metric(net, ds, dev):
27 net.eval()
28 with torch.no_grad():
29 pred = net(ds['xte'].to(dev))
30 return float(((pred - ds['yte'].to(dev)) ** 2).mean().cpu())
31
32
33def train_standard(seed, lr, epochs=EPOCHS):
34 seed_all(seed); ds = get_dataset('dynamics', seed, NTRAIN, NTEST)
35 dev = device(); net = make_model('rnn_small', tuple(ds['xtr'].shape[1:]), 1).to(dev)
36 opt = torch.optim.Adam(net.parameters(), lr=lr)
37 lossf = nn.MSELoss(); x, y = ds['xtr'].to(dev), ds['ytr'].to(dev)
38 hist=[]
39 for ep in range(epochs):
40 net.train(); perm=torch.randperm(len(x), device=dev); total=0.
41 for j in range(0,len(x),BATCH):
42 ix=perm[j:j+BATCH]; loss=lossf(net(x[ix]),y[ix]); opt.zero_grad(); loss.backward(); opt.step(); total += float(loss)*len(ix)
43 hist.append(total/len(x))
44 return metric(net,ds,dev), {'loss':hist, 'escape_times':[], 'r':0.0, 'beta':float('nan')}
45
46
47def fit_monitor(times, r):
48 if len(times) < 8 or not r > 0: return {'n':len(times), 'r':float(r), 'beta':None, 'beta_r':None, 'ks_p':None, 'gumbel_aic':None, 'exp_aic':None}
49 z=np.asarray(times,float); mu,beta=gumbel_r.fit(z); beta=abs(float(beta))
50 llg=float(np.sum(gumbel_r.logpdf(z,mu,beta))); scale=max(float(np.mean(z)),1e-8)
51 lle=float(np.sum(-z/scale-np.log(scale)))
52 ks=kstest(z,'gumbel_r',args=(mu,beta))
53 return {'n':len(z),'r':float(r),'mu':float(mu),'beta':beta,'beta_r':float(beta*r),'ks_p':float(ks.pvalue),'gumbel_aic':float(4-2*llg),'exp_aic':float(2-2*lle)}
54
55
56def train_controller(seed, lr, delay, epochs=EPOCHS):
57 seed_all(seed); ds=get_dataset('dynamics',seed,NTRAIN,NTEST); dev=device()
58 net=make_model('rnn_small',tuple(ds['xtr'].shape[1:]),1).to(dev); opt=torch.optim.Adam(net.parameters(),lr=lr)
59 lossf=nn.MSELoss(); x,y=ds['xtr'].to(dev),ds['ytr'].to(dev); hist=[]; qhist=[]; times=[]; pending=None; start_norm=None
60 step=0; halted=0
61 for ep in range(epochs):
62 net.train(); perm=torch.randperm(len(x),device=dev); total=0.
63 for j in range(0,len(x),BATCH):
64 ix=perm[j:j+BATCH]; loss=lossf(net(x[ix]),y[ix]); opt.zero_grad(); loss.backward()
65 grads=[p.grad.detach().clone() if p.grad is not None else None for p in net.parameters()]
66 noise=np.mean([float(g.var().cpu()) for g in grads if g is not None]) + 1e-8
67 q=float(sum((p.detach()**2).sum().cpu() for p in net.parameters()).sqrt()/math.sqrt(noise))
68 qhist.append(q); start_norm = start_norm or max(q,1e-6)
69 if pending is None: pending=grads
70 if step % delay == 0:
71 # delayed update; if monitoring sees explosive displacement, skip one burst
72 recent=qhist[-min(6,len(qhist)):]
73 explosive=len(recent)>=3 and recent[-1] > 2.5*max(recent[0],1e-6)
74 if explosive:
75 times.append(step); halted += 1
76 if len(times)>=8:
77 lrisk=np.polyfit(np.arange(max(0,len(qhist)-6),len(qhist)), np.log(np.maximum(qhist[-6:],1e-8)),1)[0]
78 fit=fit_monitor(times,max(float(lrisk),0.0))
79 # Gumbel gate: terminate a burst if fit is poor or divergence has no task signal
80 if fit.get('ks_p') is not None and fit['ks_p'] < .05: pending=None; step+=1; continue
81 for p,g in zip(net.parameters(),pending):
82 if g is not None: p.grad=g
83 opt.step(); pending=None
84 total += float(loss)*len(ix); step += 1
85 hist.append(total/len(x))
86 recent=np.asarray(qhist[-8:],float)
87 r=float(max(np.polyfit(np.arange(len(recent)),np.log(np.maximum(recent,1e-8)),1)[0],0.0)) if len(recent)>=3 else 0.
88 sig=fit_monitor(times,r); sig['halted_bursts']=halted; sig['q_start']=float(start_norm or 0); sig['q_end']=float(qhist[-1] if qhist else 0)
89 return metric(net,ds,dev), {'loss':hist, 'escape_times':times, **sig}
90
91
92def main():
93 # All learning rates used by either side are in the baseline grid (parity).
94 lrs=[1e-3,3e-3,1e-2]; base_grid=[{'lr':lr} for lr in lrs]
95 def base_fn(cfg): return lambda s: train_standard(s,cfg['lr'])[0]
96 base=sweep_baseline(base_fn,base_grid)
97 idea_cfgs=[{'lr':base['best_cfg']['lr'],'delay':2},{'lr':3e-3,'delay':1},{'lr':1e-2,'delay':4}]
98 idea_runs=[]
99 for cfg in idea_cfgs:
100 r=evaluate(lambda s: train_controller(s,cfg['lr'],cfg['delay'])[0],SEEDS)
101 idea_runs.append((cfg,r))
102 cfg,idea=min(idea_runs,key=lambda z:z[1]['mean'])
103 # Re-run selected idea models to obtain behavior signatures from trained models.
104 details=[train_controller(s,cfg['lr'],cfg['delay'])[1] for s in SEEDS]
105 sigs=[d for d in details if d.get('beta_r') is not None]
106 br=[d['beta_r'] for d in sigs]
107 signature={'prediction':'beta*r approximately constant across delayed bursts','observed_n':len(br),'observed_beta_r':br,'mean_beta_r':float(np.mean(br)) if br else None,'relative_spread':float(np.std(br)/abs(np.mean(br))) if br and np.mean(br)!=0 else None,'trained_model_escape_events':int(sum(len(d['escape_times']) for d in details)),'confirmed':bool(len(br)>=3 and np.std(br)/abs(np.mean(br)) <= .30)}
108 report=make_report('dynamics','rnn_small',base,idea,{'mechanism_signature':signature,'idea_sweep':[{'cfg':c,'mean':r['mean']} for c,r in idea_runs]})
109 report['selected_idea_cfg']=cfg; report['seed_details']=details
110 def clean(v):
111 if isinstance(v, float) and not math.isfinite(v): return None
112 if isinstance(v, dict): return {k: clean(x) for k,x in v.items()}
113 if isinstance(v, list): return [clean(x) for x in v]
114 return v
115 report=clean(report)
116 with open('bench_report.json','w') as f: json.dump(report,f,indent=2,allow_nan=False)
117 print(json.dumps(report,indent=2))
118
119if __name__=='__main__': main()