import sys, json, random import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report SEEDS = tuple(range(8)) GRID = [{'lr': 1e-3, 'weight_decay': 0.0}, {'lr': 3e-3, 'weight_decay': 0.0}, {'lr': 6e-3, 'weight_decay': 0.0}] EPOCHS, BATCH = 18, 64 def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def train(seed, cfg, idea=False, collect=False): seed_all(seed) ds = get_dataset('tabular', seed, n_train=400, n_test=400) net = make_model('mlp_tiny', ds['input_shape'], ds['out_dim']) dev = 'cuda' if torch.cuda.is_available() else 'cpu' try: net = net.to(dev) x, y = ds['xtr'].to(dev), ds['ytr'].to(dev) xt, yt = ds['xte'].to(dev), ds['yte'].to(dev) opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay']) lossf = nn.MSELoss() # Fixed a-priori OU persistence and observation/process variances. q, obs_var, proc_var = np.exp(-1.0 / 5.0), 0.25, 0.75 ahat = [torch.zeros_like(p) for p in net.parameters()] cov = [torch.ones_like(p) for p in net.parameters()] residual_ac, corrected_ac = [], [] prev_r, prev_c = None, None for _ in range(EPOCHS): perm = torch.randperm(len(x), device=dev) for ix in perm.split(BATCH): opt.zero_grad(set_to_none=True) lossf(net(x[ix]), y[ix]).backward() raw = [] for p in net.parameters(): raw.append(p.grad.detach().clone()) if idea: for j, p in enumerate(net.parameters()): # Position-only local prediction: the prior disturbance is # OU-persisted; the observed gradient is its measurement. hp = q * ahat[j] pp = q*q * cov[j] + proc_var gain = pp / (pp + obs_var) ahat[j] = hp + gain * (raw[j] - hp) cov[j] = (1.0 - gain) * pp p.grad.copy_(raw[j] - 0.85 * ahat[j]) cur = [raw[j] - 0.85 * ahat[j] for j in range(len(raw))] else: cur = raw if collect and prev_r is not None: a = torch.cat([z.flatten() for z in raw]).detach().float() b = torch.cat([z.flatten() for z in cur]).detach().float() u = torch.cat([z.flatten() for z in prev_r]).detach().float() v = torch.cat([z.flatten() for z in prev_c]).detach().float() def corr(z, w): z=z-z.mean(); w=w-w.mean() return float((z*w).mean()/(z.square().mean().sqrt()*w.square().mean().sqrt()+1e-8)) residual_ac.append(corr(a,u)); corrected_ac.append(corr(b,v)) prev_r, prev_c = raw, cur opt.step() with torch.no_grad(): metric = float(lossf(net(xt), yt).cpu()) sig = {'raw_grad_lag1_corr': float(np.mean(residual_ac)) if residual_ac else float('nan'), 'corrected_grad_lag1_corr': float(np.mean(corrected_ac)) if corrected_ac else float('nan'), 'predicted': 'persistent OU residual should have positive lag-1 correlation and cancellation should reduce it'} return metric, sig, net except Exception: # Robust CPU fallback after any CUDA/runtime failure. torch.cuda.empty_cache() if torch.cuda.is_available() else None return train_cpu(seed, cfg, idea, collect) def train_cpu(seed, cfg, idea=False, collect=False): seed_all(seed) ds = get_dataset('tabular', seed, n_train=400, n_test=400) net = make_model('mlp_tiny', ds['input_shape'], ds['out_dim']).cpu() x,y,xt,yt = ds['xtr'],ds['ytr'],ds['xte'],ds['yte'] opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'],weight_decay=cfg['weight_decay']); lf=nn.MSELoss() 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()] raw_ac, cor_ac, prev, prevc = [],[],None,None for _ in range(EPOCHS): for ix in torch.randperm(len(x)).split(BATCH): opt.zero_grad(); lf(net(x[ix]),y[ix]).backward(); raw=[p.grad.detach().clone() for p in net.parameters()]; cur=raw if idea: cur=[] for j,p in enumerate(net.parameters()): 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()) if prev is not None: def co(a,b): 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)) raw_ac.append(co(raw,prev));cor_ac.append(co(cur,prevc)) prev,prevc=raw,cur;opt.step() 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 def main(): def base_fn(cfg): return lambda s: train(s,cfg,False)[0] base = sweep_baseline(base_fn, GRID, seeds=(0,1,2,3)) cfgs = GRID idea_trials=[] for cfg in cfgs: r=evaluate(lambda s: train(s,cfg,True)[0], seeds=SEEDS) idea_trials.append({'cfg':cfg,'result':r}) best=min(idea_trials,key=lambda z:z['result']['mean']) idea=best['result'] sigs=[train(s,best['cfg'],True,True)[1] for s in SEEDS] bsigs=[train(s,best['cfg'],False,True)[1] for s in SEEDS] sig={k:float(np.nanmean([z[k] for z in sigs])) for k in sigs[0] if k!='predicted'} sig['baseline_raw_grad_lag1_corr']=float(np.nanmean([z['raw_grad_lag1_corr'] for z in bsigs])) sig['confirmed']=bool(sig['raw_grad_lag1_corr']>0.02 and sig['corrected_grad_lag1_corr']