Finite-Width NNGP Covariance Stabilizer / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, math, random, sys
  2import numpy as np
  3import torch
  4from torch import nn
  5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  6import bench
  7
  8SEEDS = tuple(range(8))
  9TRACK='sequence'; MODEL='transformer_tiny'; NTR=400; NTE=200
 10# Shared search space: every idea learning rate is also a baseline configuration.
 11LR_GRID=[0.001,0.003,0.009]
 12EPOCHS=12; BATCH=64
 13
 14def seed_all(s):
 15    random.seed(s); np.random.seed(s); torch.manual_seed(s)
 16    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
 17
 18def make(seed):
 19    seed_all(seed); ds=bench.get_dataset(TRACK, seed, NTR, NTE)
 20    return ds, bench.make_model(MODEL, ds['input_shape'], ds['out_dim'])
 21
 22def target_cov(x, model):
 23    # NNGP for the first shared linear projection, pooled over sequence positions.
 24    # PyTorch Linear(1,d) has init variance 1/3; positional vectors are shared.
 25    c=1.0/3.0
 26    z=x.mean(1)
 27    K=c*(z[:,None]*z[None,:])
 28    p=model.pos[:, :x.shape[1], :].mean(1)
 29    K=K + (p @ p.T)/p.shape[-1]
 30    return K.detach()
 31
 32def pooled_cov(h):
 33    q=h.mean(1)
 34    return q @ q.T / q.shape[-1]
 35
 36def train_idea(seed, lr, lam, return_model=False):
 37    ds, net=make(seed)
 38    x,y=ds['xtr'],ds['ytr']; lossf=nn.MSELoss()
 39    # Robust CUDA fallback, matching the benchmark's allowed device policy.
 40    devices=['cuda','cpu'] if torch.cuda.is_available() else ['cpu']
 41    last=None
 42    for dev in devices:
 43        try:
 44            net=net.to(dev); xx,yy=x.to(dev),y.to(dev)
 45            opt=torch.optim.Adam(net.parameters(),lr=lr)
 46            for ep in range(EPOCHS):
 47                net.train(); perm=torch.randperm(len(xx),device=dev)
 48                for i in range(0,len(xx),BATCH):
 49                    ix=perm[i:i+BATCH]; xb,yb=xx[ix],yy[ix]
 50                    h=net.inp(xb.unsqueeze(-1))+net.pos[:,:xb.shape[1]]
 51                    pred=net.head(net.enc(h).reshape(xb.shape[0],-1))
 52                    task=lossf(pred,yb)
 53                    cov=((pooled_cov(h)-target_cov(xb,net))**2).mean()
 54                    loss=task+lam*cov
 55                    opt.zero_grad(); loss.backward(); opt.step()
 56            net.eval()
 57            with torch.no_grad():
 58                pred=net(ds['xte'].to(dev)); metric=float(((pred-ds['yte'].to(dev))**2).mean())
 59            if return_model: return net, metric, ds, dev
 60            return metric
 61        except RuntimeError as e:
 62            last=e
 63            if dev=='cuda':
 64                net=net.to('cpu'); continue
 65            raise
 66    raise last
 67
 68def train_base(seed, lr, return_model=False):
 69    ds,net=make(seed)
 70    out=bench.train_model(net,ds,epochs=EPOCHS,lr=lr,batch=BATCH,weight_decay=0.0,log=lambda _:None)
 71    if return_model: return out[0],float(out[1]),ds
 72    return float(out[1])
 73
 74def signature(base_cfg, idea_cfg):
 75    bd=[]; idd=[]
 76    for s in SEEDS:
 77        bm,bmse,ds=train_base(s,base_cfg['lr'],True)
 78        im,imse,ids,dev=train_idea(s,idea_cfg['lr'],idea_cfg['lambda_cov'],True)
 79        with torch.no_grad():
 80            xb=ds['xte'][:64].to(next(bm.parameters()).device)
 81            hb=bm.inp(xb.unsqueeze(-1))+bm.pos[:,:xb.shape[1]]
 82            kb=target_cov(xb,bm).to(hb.device)
 83            xi=ids['xte'][:64].to(dev); hi=im.inp(xi.unsqueeze(-1))+im.pos[:,:xi.shape[1]]
 84            ki=target_cov(xi,im).to(hi.device)
 85            bd.append(float(((pooled_cov(hb)-kb)**2).mean()))
 86            idd.append(float(((pooled_cov(hi)-ki)**2).mean()))
 87    predicted=1/math.sqrt(64)
 88    observed_ratio=float(np.mean(idd)/np.mean(bd))
 89    return {'quantity':'pooled input-projection covariance deviation on held-out benchmark windows',
 90            'predicted_finite_width_scale_at_width_64':predicted,
 91            'observed_baseline_mean':float(np.mean(bd)),
 92            'observed_idea_mean':float(np.mean(idd)),
 93            'observed_idea_over_baseline':observed_ratio,
 94            'confirmed': bool(observed_ratio < 0.9),
 95            'note':'The O(n^-1/2) prediction is a scale law, not an absolute covariance value; this signature tests reduction on trained systems.'}
 96
 97def main():
 98    base_grid=[{'lr':v} for v in LR_GRID]
 99    base=bench.sweep_baseline(lambda c: lambda s: train_base(s,c['lr']),base_grid,seeds=SEEDS)
100    idea_trials=[]
101    for cfg in [{'lr':base['best_cfg']['lr'],'lambda_cov':1e-4}, {'lr':0.001,'lambda_cov':1e-3}, {'lr':0.009,'lambda_cov':1e-2}]:
102        r=bench.evaluate(lambda s,c=cfg: train_idea(s,c['lr'],c['lambda_cov']),SEEDS)
103        idea_trials.append({'cfg':cfg,'result':r})
104    best=min(idea_trials,key=lambda z:z['result']['mean'])
105    rep=bench.make_report(TRACK,MODEL,base,best['result'],{'mechanism_signature':signature(base['best_cfg'],best['cfg']), 'idea_sweep':idea_trials, 'epochs':EPOCHS,'batch':BATCH})
106    with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
107    print(json.dumps(rep,indent=2))
108if __name__=='__main__': main()