PSD Spectral CNN Block / run_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 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()