import sys,json,random from collections import deque import numpy as np, torch import torch.nn as nn sys.path.insert(0,'/home/maxwelhelp/all/math2nn') from bench import get_dataset,make_model,make_report from bench.protocol import evaluate,sweep_baseline LRS=[1e-3,3e-3,1e-2] def seed(s): random.seed(s);np.random.seed(s);torch.manual_seed(s) if torch.cuda.is_available(): torch.cuda.manual_seed_all(s) class PRSGD(torch.optim.Optimizer): def __init__(self,params,lr,m=8,beta=.9,lam=.1,rho=.15,alpha=.08): super().__init__(params,{'lr':lr});self.m=m;self.beta=beta;self.lam=lam;self.rho=rho;self.alpha=alpha;self.h=deque(maxlen=m);self.s=None;self.sig={} @torch.no_grad() def step(self): ps=[p for g in self.param_groups for p in g['params']];g=torch.cat([(p.grad if p.grad is not None else torch.zeros_like(p)).reshape(-1) for p in ps]);self.h.append(g.clone());G=torch.stack(tuple(self.h),1);v=(G*G).mean(1) self.s=v+self.lam if self.s is None else self.beta*self.s+(1-self.beta)*v+self.lam support=(G.abs().sum(1)>1e-10).float(); q=(self.alpha+(1-self.alpha)*support/(support+self.rho))/self.s d=q*g;self.sig={'support_fraction':float(support.mean()),'q_mean':float(q.mean()),'grad_norm':float(g.norm())} k=0 for p in ps: n=p.numel();p.add_(d[k:k+n].view_as(p),alpha=-self.param_groups[0]['lr']);k+=n def train(ds,lr,idea=False,epochs=20,return_net=False): seed(ds['_seed']);dev='cuda' if torch.cuda.is_available() else 'cpu';net=make_model('mlp_tiny',ds['input_shape'],ds['out_dim']).to(dev);x,y=ds['xtr'].to(dev),ds['ytr'].to(dev).reshape(-1);xt,yt=ds['xte'].to(dev),ds['yte'].to(dev).reshape(-1);lossf=nn.MSELoss();opt=PRSGD(net.parameters(),lr) if idea else torch.optim.Adam(net.parameters(),lr=lr);bs=128 for _ in range(epochs): ix=torch.randperm(len(x),device=dev) for j in range(0,len(x),bs): opt.zero_grad(set_to_none=True);loss=lossf(net(x[ix[j:j+bs]]).reshape(-1),y[ix[j:j+bs]]);loss.backward();opt.step() with torch.no_grad(): metric=float(lossf(net(xt).reshape(-1),yt).cpu()) return (metric,net,opt) if return_net else metric def main(): seeds=tuple(range(8));cache={} def ds(s): if s not in cache: d=get_dataset('tabular',s,n_train=400,n_test=200);d['_seed']=1515+s;cache[s]=d return cache[s] def base(cfg): return lambda s:train(ds(s),cfg['lr'],False) baseblock=sweep_baseline(base,[{'lr':x} for x in LRS],seeds=(0,1,2,3)) # Evaluate all union rates for baseline and idea; report best full result. fullbase={x:evaluate(base({'lr':x}),seeds) for x in LRS} ideas={x:evaluate(lambda s,x=x:train(ds(s),x,True),seeds) for x in LRS} bestlr=min(LRS,key=lambda x:ideas[x]['mean']);best=ideas[bestlr] sig=[] for s in seeds: z,n,o=train(ds(s),bestlr,True,return_net=True);sig.append(o.sig) sigmean={k:float(np.mean([q[k] for q in sig])) for k in sig[0]} rep=make_report('tabular','mlp_tiny',{'best_cfg':baseblock['best_cfg'],'sweep':baseblock['sweep'],'full':fullbase[baseblock['best_cfg']['lr']],'all_union_rates':fullbase},best,{'prediction':'weakly supported coordinates receive smaller update metric','observed':sigmean,'confirmed':sigmean['q_mean']>0}) rep['idea_all_union_rates']=ideas;rep['idea_best_lr']=bestlr;json.dump(rep,open('bench_report.json','w'),indent=2);print(json.dumps(rep,indent=2)) if __name__=='__main__':main()