Fourier-Calibrated Nonlocal Feature Gradient / experiment.py
Mechanism failed
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()