import sys,json,copy import numpy as np import torch import torch.nn as nn sys.path.insert(0,'/home/maxwelhelp/all/math2nn') from bench import get_dataset,train_model,make_report from bench.protocol import evaluate,sweep_baseline SEEDS=tuple(range(8)); SWEEP_SEEDS=tuple(range(4)); EPOCHS=4; NTRAIN=400; NTEST=200 GRID=[{'lr':1e-3},{'lr':3e-3},{'lr':1e-2}] class AW(nn.Module): def __init__(self,c=32,eps=1e-4,m=.1,rho=.8): super().__init__(); self.c=c; self.eps=eps; self.m=m; self.rho=rho 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)) self.observed_cov=None; self.observed_fidelity=None def forward(self,x): 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]) if self.training: 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 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 w,u=torch.linalg.eigh(use+self.eps*torch.eye(self.c,device=x.device)); a=(u*torch.rsqrt(w).unsqueeze(0))@u.T if not self.training: 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() return ((z-mu.view(1,1,1,-1))/sd.view(1,1,1,-1)@a).permute(0,3,1,2) def cnn(out=10,idea=False): class Net(nn.Module): def __init__(self): 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 def forward(self,x): 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) except RuntimeError: if self.no: raise self.no=True; old=torch.backends.cudnn.enabled; torch.backends.cudnn.enabled=False try: return self.forward(x) finally: torch.backends.cudnn.enabled=old return Net() def run(cfg,seed,idea): 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 def metric_fn(cfg,idea): return lambda seed: run(cfg,seed,idea)[0] def main(): base=sweep_baseline(lambda c:metric_fn(c,False),GRID,seeds=SWEEP_SEEDS) # Explicitly evaluate the same union/grid for idea; best chosen on sweep seeds. it=[] for c in GRID: it.append({'cfg':c,'mean':evaluate(metric_fn(c,True),SWEEP_SEEDS)['mean']}) ib=min(it,key=lambda x:x['mean'])['cfg']; idea=evaluate(metric_fn(ib,True),SEEDS) # Signature from trained systems on one paired test model, not an analytic identity. _,net,d=run(ib,0,True); net.eval(); net=net.cpu(); torch.backends.cudnn.enabled=False with torch.no_grad(): _=net(d['xte'][:128]) 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) 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()) 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)} 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'}) print(json.dumps(rep,indent=2)) if __name__=='__main__': main()