Projection-Regularized Gradient Updates / bench_projection.py

Failed on benchmark

Raw ⬇ ZIP
 1import sys,json,random
 2from collections import deque
 3import numpy as np, torch
 4import torch.nn as nn
 5sys.path.insert(0,'/home/maxwelhelp/all/math2nn')
 6from bench import get_dataset,make_model,make_report
 7from bench.protocol import evaluate,sweep_baseline
 8LRS=[1e-3,3e-3,1e-2]
 9def seed(s):
10 random.seed(s);np.random.seed(s);torch.manual_seed(s)
11 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
12class PRSGD(torch.optim.Optimizer):
13 def __init__(self,params,lr,m=8,beta=.9,lam=.1,rho=.15,alpha=.08):
14  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={}
15 @torch.no_grad()
16 def step(self):
17  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)
18  self.s=v+self.lam if self.s is None else self.beta*self.s+(1-self.beta)*v+self.lam
19  support=(G.abs().sum(1)>1e-10).float(); q=(self.alpha+(1-self.alpha)*support/(support+self.rho))/self.s
20  d=q*g;self.sig={'support_fraction':float(support.mean()),'q_mean':float(q.mean()),'grad_norm':float(g.norm())}
21  k=0
22  for p in ps:
23   n=p.numel();p.add_(d[k:k+n].view_as(p),alpha=-self.param_groups[0]['lr']);k+=n
24
25def train(ds,lr,idea=False,epochs=20,return_net=False):
26 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
27 for _ in range(epochs):
28  ix=torch.randperm(len(x),device=dev)
29  for j in range(0,len(x),bs):
30   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()
31 with torch.no_grad(): metric=float(lossf(net(xt).reshape(-1),yt).cpu())
32 return (metric,net,opt) if return_net else metric
33
34def main():
35 seeds=tuple(range(8));cache={}
36 def ds(s):
37  if s not in cache:
38   d=get_dataset('tabular',s,n_train=400,n_test=200);d['_seed']=1515+s;cache[s]=d
39  return cache[s]
40 def base(cfg): return lambda s:train(ds(s),cfg['lr'],False)
41 baseblock=sweep_baseline(base,[{'lr':x} for x in LRS],seeds=(0,1,2,3))
42 # Evaluate all union rates for baseline and idea; report best full result.
43 fullbase={x:evaluate(base({'lr':x}),seeds) for x in LRS}
44 ideas={x:evaluate(lambda s,x=x:train(ds(s),x,True),seeds) for x in LRS}
45 bestlr=min(LRS,key=lambda x:ideas[x]['mean']);best=ideas[bestlr]
46 sig=[]
47 for s in seeds:
48  z,n,o=train(ds(s),bestlr,True,return_net=True);sig.append(o.sig)
49 sigmean={k:float(np.mean([q[k] for q in sig])) for k in sig[0]}
50 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})
51 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))
52if __name__=='__main__':main()