Compactified Burst Controller / compactified_bench.py

Mechanism confirmed, baseline not beaten

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