Position-only active-noise optimizer / bench_active_noise.py
Mechanism confirmed, baseline not beaten
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()