Gumbel escape-time controller / bench_gumbel_controller.py

Failed on benchmark

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