Position-only active-noise optimizer / bench_active_noise.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, json, random
  2import numpy as np
  3import torch
  4import torch.nn as nn
  5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  6from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report
  7
  8SEEDS = tuple(range(8))
  9GRID = [{'lr': 1e-3, 'weight_decay': 0.0},
 10        {'lr': 3e-3, 'weight_decay': 0.0},
 11        {'lr': 6e-3, 'weight_decay': 0.0}]
 12EPOCHS, BATCH = 18, 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(): torch.cuda.manual_seed_all(seed)
 18
 19
 20def train(seed, cfg, idea=False, collect=False):
 21    seed_all(seed)
 22    ds = get_dataset('tabular', seed, n_train=400, n_test=400)
 23    net = make_model('mlp_tiny', ds['input_shape'], ds['out_dim'])
 24    dev = 'cuda' if torch.cuda.is_available() else 'cpu'
 25    try:
 26        net = net.to(dev)
 27        x, y = ds['xtr'].to(dev), ds['ytr'].to(dev)
 28        xt, yt = ds['xte'].to(dev), ds['yte'].to(dev)
 29        opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay'])
 30        lossf = nn.MSELoss()
 31        # Fixed a-priori OU persistence and observation/process variances.
 32        q, obs_var, proc_var = np.exp(-1.0 / 5.0), 0.25, 0.75
 33        ahat = [torch.zeros_like(p) for p in net.parameters()]
 34        cov = [torch.ones_like(p) for p in net.parameters()]
 35        residual_ac, corrected_ac = [], []
 36        prev_r, prev_c = None, None
 37        for _ in range(EPOCHS):
 38            perm = torch.randperm(len(x), device=dev)
 39            for ix in perm.split(BATCH):
 40                opt.zero_grad(set_to_none=True)
 41                lossf(net(x[ix]), y[ix]).backward()
 42                raw = []
 43                for p in net.parameters():
 44                    raw.append(p.grad.detach().clone())
 45                if idea:
 46                    for j, p in enumerate(net.parameters()):
 47                        # Position-only local prediction: the prior disturbance is
 48                        # OU-persisted; the observed gradient is its measurement.
 49                        hp = q * ahat[j]
 50                        pp = q*q * cov[j] + proc_var
 51                        gain = pp / (pp + obs_var)
 52                        ahat[j] = hp + gain * (raw[j] - hp)
 53                        cov[j] = (1.0 - gain) * pp
 54                        p.grad.copy_(raw[j] - 0.85 * ahat[j])
 55                    cur = [raw[j] - 0.85 * ahat[j] for j in range(len(raw))]
 56                else:
 57                    cur = raw
 58                if collect and prev_r is not None:
 59                    a = torch.cat([z.flatten() for z in raw]).detach().float()
 60                    b = torch.cat([z.flatten() for z in cur]).detach().float()
 61                    u = torch.cat([z.flatten() for z in prev_r]).detach().float()
 62                    v = torch.cat([z.flatten() for z in prev_c]).detach().float()
 63                    def corr(z, w):
 64                        z=z-z.mean(); w=w-w.mean()
 65                        return float((z*w).mean()/(z.square().mean().sqrt()*w.square().mean().sqrt()+1e-8))
 66                    residual_ac.append(corr(a,u)); corrected_ac.append(corr(b,v))
 67                prev_r, prev_c = raw, cur
 68                opt.step()
 69        with torch.no_grad(): metric = float(lossf(net(xt), yt).cpu())
 70        sig = {'raw_grad_lag1_corr': float(np.mean(residual_ac)) if residual_ac else float('nan'),
 71               'corrected_grad_lag1_corr': float(np.mean(corrected_ac)) if corrected_ac else float('nan'),
 72               'predicted': 'persistent OU residual should have positive lag-1 correlation and cancellation should reduce it'}
 73        return metric, sig, net
 74    except Exception:
 75        # Robust CPU fallback after any CUDA/runtime failure.
 76        torch.cuda.empty_cache() if torch.cuda.is_available() else None
 77        return train_cpu(seed, cfg, idea, collect)
 78
 79
 80def train_cpu(seed, cfg, idea=False, collect=False):
 81    seed_all(seed)
 82    ds = get_dataset('tabular', seed, n_train=400, n_test=400)
 83    net = make_model('mlp_tiny', ds['input_shape'], ds['out_dim']).cpu()
 84    x,y,xt,yt = ds['xtr'],ds['ytr'],ds['xte'],ds['yte']
 85    opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'],weight_decay=cfg['weight_decay']); lf=nn.MSELoss()
 86    q, ahat, cov = np.exp(-1/5), [torch.zeros_like(p) for p in net.parameters()], [torch.ones_like(p) for p in net.parameters()]
 87    raw_ac, cor_ac, prev, prevc = [],[],None,None
 88    for _ in range(EPOCHS):
 89        for ix in torch.randperm(len(x)).split(BATCH):
 90            opt.zero_grad(); lf(net(x[ix]),y[ix]).backward(); raw=[p.grad.detach().clone() for p in net.parameters()]; cur=raw
 91            if idea:
 92                cur=[]
 93                for j,p in enumerate(net.parameters()):
 94                    hp=q*ahat[j]; pp=q*q*cov[j]+.75; k=pp/(pp+.25); ahat[j]=hp+k*(raw[j]-hp); cov[j]=(1-k)*pp; p.grad.copy_(raw[j]-.85*ahat[j]); cur.append(p.grad.detach().clone())
 95            if prev is not None:
 96                def co(a,b):
 97                    a=torch.cat([z.flatten() for z in a]);b=torch.cat([z.flatten() for z in b]);a-=a.mean();b-=b.mean();return float((a*b).mean()/(a.square().mean().sqrt()*b.square().mean().sqrt()+1e-8))
 98                raw_ac.append(co(raw,prev));cor_ac.append(co(cur,prevc))
 99            prev,prevc=raw,cur;opt.step()
100    return float(lf(net(xt),yt)), {'raw_grad_lag1_corr':float(np.mean(raw_ac)),'corrected_grad_lag1_corr':float(np.mean(cor_ac))},net
101
102
103def main():
104    def base_fn(cfg): return lambda s: train(s,cfg,False)[0]
105    base = sweep_baseline(base_fn, GRID, seeds=(0,1,2,3))
106    cfgs = GRID
107    idea_trials=[]
108    for cfg in cfgs:
109        r=evaluate(lambda s: train(s,cfg,True)[0], seeds=SEEDS)
110        idea_trials.append({'cfg':cfg,'result':r})
111    best=min(idea_trials,key=lambda z:z['result']['mean'])
112    idea=best['result']
113    sigs=[train(s,best['cfg'],True,True)[1] for s in SEEDS]
114    bsigs=[train(s,best['cfg'],False,True)[1] for s in SEEDS]
115    sig={k:float(np.nanmean([z[k] for z in sigs])) for k in sigs[0] if k!='predicted'}
116    sig['baseline_raw_grad_lag1_corr']=float(np.nanmean([z['raw_grad_lag1_corr'] for z in bsigs]))
117    sig['confirmed']=bool(sig['raw_grad_lag1_corr']>0.02 and sig['corrected_grad_lag1_corr']<sig['raw_grad_lag1_corr'])
118    rep=make_report('tabular','mlp_tiny',base,idea,{'mechanism_signature':sig,'idea_trials':idea_trials,'task_match':'optimizer intervention on Friedman regression'})
119    open('bench_report.json','w').write(json.dumps(rep,indent=2))
120    print(json.dumps(rep,indent=2))
121
122if __name__=='__main__': main()