Fourier-Calibrated Nonlocal Feature Gradient / experiment.py

Mechanism failed

Raw ⬇ ZIP
  1import json, math, random, time
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6import torch.nn.functional as F
  7from sklearn.datasets import load_digits
  8from sklearn.model_selection import train_test_split
  9
 10SEED = 348
 11random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
 12try: torch.cuda.manual_seed_all(SEED)
 13except Exception: pass
 14
 15def device_or_cpu():
 16    if torch.cuda.is_available():
 17        try:
 18            torch.zeros(1, device='cuda'); return torch.device('cuda')
 19        except Exception: pass
 20    return torch.device('cpu')
 21
 22def offsets_and_weights(delta, sigma=None):
 23    sigma = float(sigma if sigma is not None else max(delta / 1.5, .5))
 24    qs = [(dy, dx) for dy in range(-delta, delta+1) for dx in range(-delta, delta+1)
 25          if dx*dx + dy*dy <= delta*delta and (dx != 0 or dy != 0)]
 26    q = np.asarray(qs, dtype=np.float32)
 27    w = np.exp(-np.sum(q*q, axis=1)/(2*sigma*sigma)).astype(np.float32); w /= w.sum()
 28    return q, w
 29
 30def multiplier(h, w, delta, sigma=None):
 31    q, weights = offsets_and_weights(delta, sigma)
 32    ky=np.arange(h,dtype=np.float32)[:,None]; kx=np.arange(w,dtype=np.float32)[None,:]
 33    m=np.zeros((h,w),dtype=np.float32)
 34    for (dy,dx),wt in zip(q,weights):
 35        m += wt*(1-np.cos(2*np.pi*(ky*dy/h+kx*dx/w)))
 36    return m
 37
 38def math_check():
 39    h=w=32; m1=multiplier(h,w,1); m3=multiplier(h,w,3)
 40    low=float(m3[1,0]); high=float(m3[h//2,0])
 41    rng=np.random.default_rng(SEED); z=rng.normal(size=(h,w)).astype(np.float32); zh=np.fft.fft2(z)
 42    spectral=float(np.sum(m3*np.abs(zh)**2)/(h*w)); q,weights=offsets_and_weights(3); spatial=0.
 43    for (dy,dx),wt in zip(q,weights):
 44        d=np.roll(z,(int(dy),int(dx)),axis=(0,1))-z; spatial += .5*float(wt)*float(np.sum(d*d))
 45    # This is the literal scalar-weighted G in the pseudocode. Symmetry makes it vanish.
 46    x=np.arange(w,dtype=np.float32)[None,:] + 2*np.arange(h,dtype=np.float32)[:,None]
 47    g=np.zeros_like(x)
 48    for (dy,dx),wt in zip(q,weights): g += wt*(np.roll(x,(int(dy),int(dx)),axis=(0,1))-x)
 49    return {'m_min':float(m3.min()),'m_max':float(m3.max()),'low_frequency_m':low,
 50      'nyquist_m':high,'high_to_low_ratio':high/(low+1e-12),
 51      'parseval_relative_error':abs(spectral-spatial)/(abs(spatial)+1e-12),
 52      'delta1_mean_multiplier':float(m1.mean()),'delta3_mean_multiplier':float(m3.mean()),
 53      'positive_multiplier':bool(m3.min()>=-1e-6),'high_frequency_more_penalized':bool(high>low),
 54      'literal_symmetric_G_linear_rms':float(np.sqrt(np.mean(g*g)))}
 55
 56class NonlocalBlock(nn.Module):
 57    def __init__(self,channels,delta=2):
 58        super().__init__(); self.delta=delta
 59        q,wt=offsets_and_weights(delta); self.register_buffer('q',torch.tensor(q)); self.register_buffer('wt',torch.tensor(wt))
 60        self.mix=nn.Sequential(nn.Conv2d(2*channels,channels,1),nn.BatchNorm2d(channels),nn.ReLU()); self.out=nn.Conv2d(channels,channels,1)
 61    def forward(self,x,return_penalty=False):
 62        g=torch.zeros_like(x)
 63        for i in range(self.q.shape[0]):
 64            dy,dx=int(self.q[i,0]),int(self.q[i,1]); g += self.wt[i]*(torch.roll(x,(dy,dx),dims=(2,3))-x)
 65        y=x+self.out(self.mix(torch.cat([x,g],dim=1)))
 66        if not return_penalty:return y
 67        m=torch.tensor(multiplier(x.shape[-2],x.shape[-1],self.delta),device=x.device)
 68        f=torch.fft.rfft2(y); penalty=(m[:,:f.shape[-1]][None,None]*f.abs().square()).mean()/(x.shape[-2]*x.shape[-1])
 69        return y,penalty
 70
 71class Net(nn.Module):
 72    def __init__(self,kind='baseline',delta=2):
 73        super().__init__(); self.kind=kind; self.c1=nn.Conv2d(1,16,3,padding=1); self.c2=nn.Conv2d(16,32,3,padding=1)
 74        self.nl=NonlocalBlock(16,delta) if kind=='nonlocal' else None; self.head=nn.Linear(32*2*2,10)
 75    def forward(self,x):
 76        x=F.relu(self.c1(x)); pen=x.new_zeros(())
 77        if self.nl is not None:x,pen=self.nl(x,True)
 78        x=F.max_pool2d(x,2); x=F.relu(self.c2(x)); x=F.max_pool2d(x,2); return self.head(x.flatten(1)),pen
 79
 80def run_training(kind,device,epochs=12):
 81    data=load_digits(); X=data.images.astype('float32')/16.; y=data.target
 82    Xtr,Xte,ytr,yte=train_test_split(X,y,test_size=.25,random_state=SEED,stratify=y)
 83    trX=torch.tensor(Xtr[:,None],device=device); tey=torch.tensor(ytr,device=device); vaX=torch.tensor(Xte[:,None],device=device); vay=torch.tensor(yte,device=device)
 84    torch.manual_seed(SEED); model=Net(kind).to(device); opt=torch.optim.Adam(model.parameters(),lr=2e-3); t0=time.perf_counter(); losses=[]; gs=[]
 85    for _ in range(epochs):
 86        model.train(); opt.zero_grad(set_to_none=True); logits,pen=model(trX); loss=F.cross_entropy(logits,tey)+(1e-5*pen if kind=='nonlocal' else 0.); loss.backward(); gs.append(float(torch.nn.utils.clip_grad_norm_(model.parameters(),1e9).detach().cpu())); opt.step(); losses.append(float(loss.detach().cpu()))
 87    model.eval()
 88    with torch.no_grad():
 89        pred,_=model(vaX); acc=float((pred.argmax(1)==vay).float().mean().cpu()); feat=F.relu(model.c1(vaX));
 90        if model.nl is not None: feat,_=model.nl(feat,True)
 91        fh=torch.fft.rfft2(feat); hf=fh[:,:,2:,:].abs().square().mean().item()
 92    return {'accuracy':acc,'final_loss':losses[-1],'grad_norm_std':float(np.std(gs)),'output_high_frequency_energy':hf,'seconds':time.perf_counter()-t0}
 93
 94def main():
 95    device=device_or_cpu(); result={'device':str(device),'math_check':math_check()}
 96    try:
 97        result['baseline']=run_training('baseline',device); result['nonlocal']=run_training('nonlocal',device)
 98    except Exception as exc:
 99        if device.type!='cuda': raise
100        result['cuda_error']=repr(exc); device=torch.device('cpu'); result['device']='cpu (CUDA runtime fallback)'; result['baseline']=run_training('baseline',device); result['nonlocal']=run_training('nonlocal',device)
101    Path('results.json').write_text(json.dumps(result,indent=2)); print(json.dumps(result,indent=2))
102if __name__=='__main__': main()