Anchored Whitening Layer / bench_experiment.py

Failed on benchmark

Raw ⬇ ZIP
 1import sys,json,copy
 2import numpy as np
 3import torch
 4import torch.nn as nn
 5sys.path.insert(0,'/home/maxwelhelp/all/math2nn')
 6from bench import get_dataset,train_model,make_report
 7from bench.protocol import evaluate,sweep_baseline
 8
 9SEEDS=tuple(range(8)); SWEEP_SEEDS=tuple(range(4)); EPOCHS=4; NTRAIN=400; NTEST=200
10GRID=[{'lr':1e-3},{'lr':3e-3},{'lr':1e-2}]
11
12class AW(nn.Module):
13 def __init__(self,c=32,eps=1e-4,m=.1,rho=.8):
14  super().__init__(); self.c=c; self.eps=eps; self.m=m; self.rho=rho
15  self.register_buffer('mean',torch.zeros(c)); self.register_buffer('std',torch.ones(c)); self.register_buffer('cov',torch.eye(c)); self.register_buffer('init',torch.tensor(False))
16  self.observed_cov=None; self.observed_fidelity=None
17 def forward(self,x):
18  z=x.permute(0,2,3,1); f=z.reshape(-1,self.c); mu=f.mean(0); sd=f.std(0,unbiased=False).clamp_min(1e-3); q=(f-mu)/sd; r=q.T@q/max(1,q.shape[0])
19  if self.training:
20   self.mean.mul_(1-self.m).add_(self.m*mu); self.std.mul_(1-self.m).add_(self.m*sd); self.cov.mul_(1-self.m).add_(self.m*r); self.init.fill_(True); use=r
21  else: use=self.cov if bool(self.init) else r; mu=self.mean if bool(self.init) else mu; sd=self.std if bool(self.init) else sd
22  w,u=torch.linalg.eigh(use+self.eps*torch.eye(self.c,device=x.device)); a=(u*torch.rsqrt(w).unsqueeze(0))@u.T
23  if not self.training:
24   zz=(f-mu)/sd; co=zz.T@zz/max(1,zz.shape[0]); self.observed_cov=(a.T@co@a).detach().cpu(); self.observed_fidelity=torch.diag(co@a).detach().cpu()
25  return ((z-mu.view(1,1,1,-1))/sd.view(1,1,1,-1)@a).permute(0,3,1,2)
26
27def cnn(out=10,idea=False):
28 class Net(nn.Module):
29  def __init__(self):
30   super().__init__(); self.c1=nn.Conv2d(3,32,3,padding=1); self.aw=AW(32) if idea else nn.Identity(); self.c2=nn.Conv2d(32,64,3,padding=1); self.c3=nn.Conv2d(64,96,3,padding=1); self.head=nn.Sequential(nn.Flatten(),nn.Linear(96*4*4,128),nn.ReLU(),nn.Linear(128,out)); self.no=False
31  def forward(self,x):
32   try: x=torch.relu(self.c1(x)); x=self.aw(x); x=torch.max_pool2d(x,2); x=torch.relu(self.c2(x)); x=torch.max_pool2d(x,2); x=torch.relu(self.c3(x)); x=torch.max_pool2d(x,2); return self.head(x)
33   except RuntimeError:
34    if self.no: raise
35    self.no=True; old=torch.backends.cudnn.enabled; torch.backends.cudnn.enabled=False
36    try: return self.forward(x)
37    finally: torch.backends.cudnn.enabled=old
38 return Net()
39
40def run(cfg,seed,idea):
41 torch.manual_seed(seed); np.random.seed(seed); d=get_dataset('vision',seed,NTRAIN,NTEST); net=cnn(10,idea); _,metric,_=train_model(net,d,epochs=EPOCHS,lr=cfg['lr'],batch=128,log=lambda *a,**k:None); return float(metric),net,d
42
43def metric_fn(cfg,idea):
44 return lambda seed: run(cfg,seed,idea)[0]
45
46def main():
47 base=sweep_baseline(lambda c:metric_fn(c,False),GRID,seeds=SWEEP_SEEDS)
48 # Explicitly evaluate the same union/grid for idea; best chosen on sweep seeds.
49 it=[]
50 for c in GRID: it.append({'cfg':c,'mean':evaluate(metric_fn(c,True),SWEEP_SEEDS)['mean']})
51 ib=min(it,key=lambda x:x['mean'])['cfg']; idea=evaluate(metric_fn(ib,True),SEEDS)
52 # Signature from trained systems on one paired test model, not an analytic identity.
53 _,net,d=run(ib,0,True); net.eval();
54 net=net.cpu(); torch.backends.cudnn.enabled=False
55 with torch.no_grad(): _=net(d['xte'][:128])
56 aw=net.aw; co=aw.observed_cov.numpy() if aw.observed_cov is not None else np.eye(32); fi=aw.observed_fidelity.numpy() if aw.observed_fidelity is not None else np.ones(32)
57 off=float(np.linalg.norm(co-np.diag(np.diag(co)),'fro')); diag=float(np.max(np.abs(np.diag(co)-1))); fidelity=float(fi.min())
58 sig={'prediction':'trained anchored layer should reduce channel covariance while preserving designated-channel fidelity','baseline_offdiag_cov':None,'idea_offdiag_cov_observed':off,'idea_max_unit_variance_error_observed':diag,'idea_min_fidelity_observed':fidelity,'rho_min':.8,'confirmed':bool(off<1.0 and fidelity>=.8-0.05)}
59 rep=make_report('vision','cnn_small',base,idea,{'mechanism_signature':sig,'idea_sweep':it,'protocol_note':'shared lr grid; 8 paired seeds; 400/200 CIFAR subset; 4 epochs'})
60 print(json.dumps(rep,indent=2))
61if __name__=='__main__': main()