import json, math, random, time from pathlib import Path import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from sklearn.datasets import load_digits from sklearn.model_selection import train_test_split SEED = 348 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) try: torch.cuda.manual_seed_all(SEED) except Exception: pass def device_or_cpu(): if torch.cuda.is_available(): try: torch.zeros(1, device='cuda'); return torch.device('cuda') except Exception: pass return torch.device('cpu') def offsets_and_weights(delta, sigma=None): sigma = float(sigma if sigma is not None else max(delta / 1.5, .5)) qs = [(dy, dx) for dy in range(-delta, delta+1) for dx in range(-delta, delta+1) if dx*dx + dy*dy <= delta*delta and (dx != 0 or dy != 0)] q = np.asarray(qs, dtype=np.float32) w = np.exp(-np.sum(q*q, axis=1)/(2*sigma*sigma)).astype(np.float32); w /= w.sum() return q, w def multiplier(h, w, delta, sigma=None): q, weights = offsets_and_weights(delta, sigma) ky=np.arange(h,dtype=np.float32)[:,None]; kx=np.arange(w,dtype=np.float32)[None,:] m=np.zeros((h,w),dtype=np.float32) for (dy,dx),wt in zip(q,weights): m += wt*(1-np.cos(2*np.pi*(ky*dy/h+kx*dx/w))) return m def math_check(): h=w=32; m1=multiplier(h,w,1); m3=multiplier(h,w,3) low=float(m3[1,0]); high=float(m3[h//2,0]) rng=np.random.default_rng(SEED); z=rng.normal(size=(h,w)).astype(np.float32); zh=np.fft.fft2(z) spectral=float(np.sum(m3*np.abs(zh)**2)/(h*w)); q,weights=offsets_and_weights(3); spatial=0. for (dy,dx),wt in zip(q,weights): d=np.roll(z,(int(dy),int(dx)),axis=(0,1))-z; spatial += .5*float(wt)*float(np.sum(d*d)) # This is the literal scalar-weighted G in the pseudocode. Symmetry makes it vanish. x=np.arange(w,dtype=np.float32)[None,:] + 2*np.arange(h,dtype=np.float32)[:,None] g=np.zeros_like(x) for (dy,dx),wt in zip(q,weights): g += wt*(np.roll(x,(int(dy),int(dx)),axis=(0,1))-x) return {'m_min':float(m3.min()),'m_max':float(m3.max()),'low_frequency_m':low, 'nyquist_m':high,'high_to_low_ratio':high/(low+1e-12), 'parseval_relative_error':abs(spectral-spatial)/(abs(spatial)+1e-12), 'delta1_mean_multiplier':float(m1.mean()),'delta3_mean_multiplier':float(m3.mean()), 'positive_multiplier':bool(m3.min()>=-1e-6),'high_frequency_more_penalized':bool(high>low), 'literal_symmetric_G_linear_rms':float(np.sqrt(np.mean(g*g)))} class NonlocalBlock(nn.Module): def __init__(self,channels,delta=2): super().__init__(); self.delta=delta q,wt=offsets_and_weights(delta); self.register_buffer('q',torch.tensor(q)); self.register_buffer('wt',torch.tensor(wt)) self.mix=nn.Sequential(nn.Conv2d(2*channels,channels,1),nn.BatchNorm2d(channels),nn.ReLU()); self.out=nn.Conv2d(channels,channels,1) def forward(self,x,return_penalty=False): g=torch.zeros_like(x) for i in range(self.q.shape[0]): 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) y=x+self.out(self.mix(torch.cat([x,g],dim=1))) if not return_penalty:return y m=torch.tensor(multiplier(x.shape[-2],x.shape[-1],self.delta),device=x.device) f=torch.fft.rfft2(y); penalty=(m[:,:f.shape[-1]][None,None]*f.abs().square()).mean()/(x.shape[-2]*x.shape[-1]) return y,penalty class Net(nn.Module): def __init__(self,kind='baseline',delta=2): super().__init__(); self.kind=kind; self.c1=nn.Conv2d(1,16,3,padding=1); self.c2=nn.Conv2d(16,32,3,padding=1) self.nl=NonlocalBlock(16,delta) if kind=='nonlocal' else None; self.head=nn.Linear(32*2*2,10) def forward(self,x): x=F.relu(self.c1(x)); pen=x.new_zeros(()) if self.nl is not None:x,pen=self.nl(x,True) 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 def run_training(kind,device,epochs=12): data=load_digits(); X=data.images.astype('float32')/16.; y=data.target Xtr,Xte,ytr,yte=train_test_split(X,y,test_size=.25,random_state=SEED,stratify=y) 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) torch.manual_seed(SEED); model=Net(kind).to(device); opt=torch.optim.Adam(model.parameters(),lr=2e-3); t0=time.perf_counter(); losses=[]; gs=[] for _ in range(epochs): 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())) model.eval() with torch.no_grad(): pred,_=model(vaX); acc=float((pred.argmax(1)==vay).float().mean().cpu()); feat=F.relu(model.c1(vaX)); if model.nl is not None: feat,_=model.nl(feat,True) fh=torch.fft.rfft2(feat); hf=fh[:,:,2:,:].abs().square().mean().item() return {'accuracy':acc,'final_loss':losses[-1],'grad_norm_std':float(np.std(gs)),'output_high_frequency_energy':hf,'seconds':time.perf_counter()-t0} def main(): device=device_or_cpu(); result={'device':str(device),'math_check':math_check()} try: result['baseline']=run_training('baseline',device); result['nonlocal']=run_training('nonlocal',device) except Exception as exc: if device.type!='cuda': raise 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) Path('results.json').write_text(json.dumps(result,indent=2)); print(json.dumps(result,indent=2)) if __name__=='__main__': main()