import json, random import numpy as np import torch import torch.nn as nn SEED=89 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) try: device=torch.device('cuda' if torch.cuda.is_available() else 'cpu') if device.type=='cuda': torch.zeros(1,device=device) except Exception: device=torch.device('cpu') def conv(x,k): y=torch.zeros(x.size(0),k.size(0),x.size(2),x.size(3),device=x.device) for a in range(k.size(2)): for b in range(k.size(3)): y += torch.einsum('oi,bihw->bohw',k[:,:,a,b],torch.roll(x,(-a,-b),(2,3))) return y def adj(x,k): y=torch.zeros(x.size(0),k.size(1),x.size(2),x.size(3),device=x.device) for a in range(k.size(2)): for b in range(k.size(3)): y += torch.einsum('oi,bohw->bihw',k[:,:,a,b],torch.roll(x,(a,b),(2,3))) return y def spectrum(k,H,W): z=k.detach().cpu().numpy(); vals=[] for p in range(H): for q in range(W): phase=np.exp(2j*np.pi*(np.arange(z.shape[2])[:,None]*p/H+np.arange(z.shape[3])[None,:]*q/W)) B=(z*phase[None,None]).sum((2,3)) ev=np.linalg.eigvalsh(B.conj().T@B) vals.append((ev.min(),ev.max())) return np.asarray(vals) class PSD(nn.Module): def __init__(self,n,m,e): super().__init__(); self.B=nn.Parameter(.18*torch.randn(m,n,e,e)) def forward(self,x,tau): return x-tau*adj(conv(x,self.B),self.B) class Free(nn.Module): def __init__(self,n,e): super().__init__(); self.K=nn.Parameter(.18*torch.randn(n,n,e,e)) def forward(self,x,tau): return x-tau*conv(x,self.K) def fit(model, train_x, train_y, test_x, test_y, steps=300, tau=.1): opt=torch.optim.Adam(model.parameters(),lr=.03) for _ in range(steps): opt.zero_grad(); loss=((model(train_x,tau)-train_y)**2).mean(); loss.backward(); opt.step() with torch.no_grad(): tr=((model(train_x,tau)-train_y)**2).mean().item(); te=((model(test_x,tau)-test_y)**2).mean().item() return tr,te def main(): torch.set_num_threads(4); n=2; H=W=12; e=3 k=.25*torch.randn(3,n,e,e,device=device) 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) inner=(conv(x,k)*u).sum().item(); adjinner=(x*adj(u,k)).sum().item() lam=float(s[:,1].max()); tau=1.0/(lam+1e-6) ratios=[] for _ in range(20): v=torch.randn(1,n,H,W,device=device); ratios.append((v-tau*adj(conv(v,k),k)).norm().item()/v.norm().item()) # learn a known smoothing residual map from noisy examples torch.manual_seed(SEED+1) clean=torch.randn(48,n,H,W,device=device) noise=.35*torch.randn_like(clean); noisy=clean+noise tx,ty=noisy[:32],clean[:32]; vx,vy=noisy[32:],clean[32:] torch.manual_seed(SEED+2); p=PSD(n,3,e).to(device) torch.manual_seed(SEED+2); f=Free(n,e).to(device) ptr,pte=fit(p,tx,ty,vx,vy,tau=.1); ftr,fte=fit(f,tx,ty,vx,vy,tau=.1) 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())}} print(json.dumps(result,indent=2)) if __name__=='__main__': main()