PSD Spectral CNN Block / run_experiment.py
Mechanism confirmed, baseline not beaten
1import json, random
2import numpy as np
3import torch
4import torch.nn as nn
5
6SEED=89
7random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
8try:
9 device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
10 if device.type=='cuda': torch.zeros(1,device=device)
11except Exception:
12 device=torch.device('cpu')
13
14def conv(x,k):
15 y=torch.zeros(x.size(0),k.size(0),x.size(2),x.size(3),device=x.device)
16 for a in range(k.size(2)):
17 for b in range(k.size(3)):
18 y += torch.einsum('oi,bihw->bohw',k[:,:,a,b],torch.roll(x,(-a,-b),(2,3)))
19 return y
20
21def adj(x,k):
22 y=torch.zeros(x.size(0),k.size(1),x.size(2),x.size(3),device=x.device)
23 for a in range(k.size(2)):
24 for b in range(k.size(3)):
25 y += torch.einsum('oi,bohw->bihw',k[:,:,a,b],torch.roll(x,(a,b),(2,3)))
26 return y
27
28def spectrum(k,H,W):
29 z=k.detach().cpu().numpy(); vals=[]
30 for p in range(H):
31 for q in range(W):
32 phase=np.exp(2j*np.pi*(np.arange(z.shape[2])[:,None]*p/H+np.arange(z.shape[3])[None,:]*q/W))
33 B=(z*phase[None,None]).sum((2,3))
34 ev=np.linalg.eigvalsh(B.conj().T@B)
35 vals.append((ev.min(),ev.max()))
36 return np.asarray(vals)
37
38class PSD(nn.Module):
39 def __init__(self,n,m,e):
40 super().__init__(); self.B=nn.Parameter(.18*torch.randn(m,n,e,e))
41 def forward(self,x,tau): return x-tau*adj(conv(x,self.B),self.B)
42
43class Free(nn.Module):
44 def __init__(self,n,e):
45 super().__init__(); self.K=nn.Parameter(.18*torch.randn(n,n,e,e))
46 def forward(self,x,tau): return x-tau*conv(x,self.K)
47
48def fit(model, train_x, train_y, test_x, test_y, steps=300, tau=.1):
49 opt=torch.optim.Adam(model.parameters(),lr=.03)
50 for _ in range(steps):
51 opt.zero_grad(); loss=((model(train_x,tau)-train_y)**2).mean(); loss.backward(); opt.step()
52 with torch.no_grad():
53 tr=((model(train_x,tau)-train_y)**2).mean().item(); te=((model(test_x,tau)-test_y)**2).mean().item()
54 return tr,te
55
56def main():
57 torch.set_num_threads(4); n=2; H=W=12; e=3
58 k=.25*torch.randn(3,n,e,e,device=device)
59 s=spectrum(k,H,W); x=torch.randn(4,n,H,W,device=device); u=torch.randn(4,k.size(0),H,W,device=device)
60 inner=(conv(x,k)*u).sum().item(); adjinner=(x*adj(u,k)).sum().item()
61 lam=float(s[:,1].max()); tau=1.0/(lam+1e-6)
62 ratios=[]
63 for _ in range(20):
64 v=torch.randn(1,n,H,W,device=device); ratios.append((v-tau*adj(conv(v,k),k)).norm().item()/v.norm().item())
65 # learn a known smoothing residual map from noisy examples
66 torch.manual_seed(SEED+1)
67 clean=torch.randn(48,n,H,W,device=device)
68 noise=.35*torch.randn_like(clean); noisy=clean+noise
69 tx,ty=noisy[:32],clean[:32]; vx,vy=noisy[32:],clean[32:]
70 torch.manual_seed(SEED+2); p=PSD(n,3,e).to(device)
71 torch.manual_seed(SEED+2); f=Free(n,e).to(device)
72 ptr,pte=fit(p,tx,ty,vx,vy,tau=.1); ftr,fte=fit(f,tx,ty,vx,vy,tau=.1)
73 result={'device':str(device),'fourier_min_eigenvalue':float(s[:,0].min()),'fourier_max_eigenvalue':lam,'adjoint_relative_error':abs(inner-adjinner)/(abs(inner)+abs(adjinner)+1e-12),'diffusion_tau':tau,'diffusion_max_norm_ratio':max(ratios),'psd_train_mse':ptr,'psd_test_mse':pte,'free_train_mse':ftr,'free_test_mse':fte,'parameter_counts':{'psd':sum(q.numel() for q in p.parameters()),'free':sum(q.numel() for q in f.parameters())}}
74 print(json.dumps(result,indent=2))
75
76if __name__=='__main__': main()