Hypoelliptic transport-diffusion layer / experiment.py
Mechanism confirmed, baseline not beaten
1import json, math, random
2import numpy as np
3import torch
4import torch.nn as nn
5import torch.nn.functional as F
6
7SEED=2852
8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
9try:
10 device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
11except Exception: device=torch.device('cpu')
12
13# periodic y interpolation; x diffusion is a normalized Gaussian depthwise convolution
14def diffusion_x(f, h, dx):
15 # f [B,C,Nx,Ny], Gaussian variance 2h
16 sigma=math.sqrt(2*h)/dx
17 r=max(1,int(math.ceil(3*sigma)))
18 z=torch.arange(-r,r+1,device=f.device,dtype=f.dtype)
19 w=torch.exp(-0.5*(z/sigma)**2); w=w/w.sum()
20 out=torch.zeros_like(f)
21 for i,wi in enumerate(w): out += wi*torch.roll(f, int(i-r), dims=2)
22 return out
23
24def transport_y(f,h,x,ymin=-math.pi,ymax=math.pi):
25 # source f(x, y+h*x); periodic linear interpolation
26 ny=f.shape[-1]; dy=(ymax-ymin)/ny
27 shift=h*x[:,None] # [Nx,1]
28 yy=(torch.arange(ny,device=f.device,dtype=f.dtype)[None,:]*dy+ymin+shift)
29 q=(yy-ymin)/dy % ny
30 j0=torch.floor(q).long(); a=(q-j0).to(f.dtype); j1=(j0+1)%ny
31 return f.gather(3,j0[None,None].expand(f.shape[0],f.shape[1],-1,-1))*(1-a)[None,None] + f.gather(3,j1[None,None].expand(f.shape[0],f.shape[1],-1,-1))*a[None,None]
32
33def kinetic(f,h,x,dx): return transport_y(diffusion_x(f,h,dx),h,x)
34
35class KineticLayer(nn.Module):
36 def __init__(self,c,h,x,dx):
37 super().__init__(); self.h=h; self.register_buffer('x',x); self.dx=dx
38 self.gate=nn.Conv2d(c,c,1); nn.init.constant_(self.gate.weight,0); nn.init.constant_(self.gate.bias,1.5)
39 self.mix=nn.Conv2d(c,c,1)
40 def forward(self,f):
41 z=kinetic(f,self.h,self.x,self.dx)
42 g=torch.sigmoid(self.gate(f)); return self.mix(f+g*(z-f))
43class LocalBaseline(nn.Module):
44 def __init__(self,c):
45 super().__init__(); self.net=nn.Sequential(nn.Conv2d(c,c,3,padding=1,padding_mode='circular'),nn.GELU(),nn.Conv2d(c,c,1))
46 def forward(self,f): return self.net(f)
47
48def make_fields(n,c,nx,ny,device):
49 # Smooth random phase-space fields, with enough variation to expose y transport.
50 q=torch.randn(n,c,nx,ny,device=device)
51 for _ in range(3): q=(q+torch.roll(q,1,2)+torch.roll(q,-1,2)+torch.roll(q,1,3)+torch.roll(q,-1,3))/5
52 return q
53
54def math_checks():
55 nx,ny=64,128; xmin,xmax=-2,2; dx=(xmax-xmin)/nx
56 x=torch.arange(nx,dtype=torch.float64)*dx+xmin; yy=torch.arange(ny,dtype=torch.float64)*(2*math.pi/ny)-math.pi
57 # diffusion Fourier mode prediction: exp(-h*k^2)
58 k=3.0; xx=x[:,None]; f=torch.cos(k*xx).expand(1,1,nx,ny).clone()
59 hs=np.array([.01,.025,.05,.1,.2]); ratios=[]
60 for h in hs:
61 d=diffusion_x(f,h,dx); ratios.append((d.abs().mean()/f.abs().mean()).item())
62 pred=np.exp(-hs*k*k); rel=float(np.max(np.abs(np.array(ratios)-pred)/pred))
63 # transport preserves L2; convex gated update is nonexpansive for constant g against zero
64 rng=torch.Generator().manual_seed(SEED); r=torch.randn(1,1,nx,ny,generator=rng,dtype=torch.float64)
65 norms=[]; disp=[]
66 for h in [.02,.05,.1,.2]:
67 z=kinetic(r,h,x,dx); norms.append((z.norm()/r.norm()).item())
68 # characteristic displacement RMS (on physical coordinates)
69 disp.append(h*float(torch.sqrt(torch.mean(x*x))))
70 # contraction test for g=0,.25,.5,1 on a random vector; diffusion+shift operator is empirically contractive
71 contractions=[]
72 for g in [0,.25,.5,1]:
73 z=r+g*(kinetic(r,.1,x,dx)-r); contractions.append((z.norm()/r.norm()).item())
74 return {'diffusion_h':hs.tolist(),'diffusion_observed':ratios,'diffusion_predicted':pred.tolist(),'diffusion_max_relative_error':rel,
75 'kinetic_norm_ratios_h_[.02,.05,.1,.2]':norms,'transport_rms_displacement':disp,
76 'displacement_over_h':(np.array(disp)/np.array([.02,.05,.1,.2])).tolist(),
77 'gated_norm_ratios_g_[0,.25,.5,1]':contractions}
78
79def train_compare():
80 c,nx,ny=2,24,48; h=.12; xmin,xmax=-2,2; dx=(xmax-xmin)/nx
81 x=torch.arange(nx,device=device)*dx+xmin
82 train=make_fields(96,c,nx,ny,device); val=make_fields(32,c,nx,ny,device)
83 with torch.no_grad(): target=kinetic(train,h,x,dx); vtarget=kinetic(val,h,x,dx)
84 models={'baseline':LocalBaseline(c).to(device),'idea':KineticLayer(c,h,x,dx).to(device)}
85 results={}
86 for name,m in models.items():
87 opt=torch.optim.Adam(m.parameters(),lr=3e-3)
88 for step in range(350):
89 idx=torch.randint(0,len(train),(16,),device=device); out=m(train[idx]); loss=F.mse_loss(out,target[idx])
90 opt.zero_grad(); loss.backward(); opt.step()
91 with torch.no_grad():
92 pred=m(val); mse=F.mse_loss(pred,vtarget).item()
93 # rollout by repeatedly applying learned transition, measured against true operator
94 a=val.clone(); b=val.clone(); errs_a=[]; errs_b=[]
95 for _ in range(5):
96 a=m(a); b=kinetic(b,h,x,dx); errs_a.append(F.mse_loss(a,b).item())
97 results[name]={'val_one_step_mse':mse,'rollout_mse_steps_1_to_5':errs_a,'parameters':sum(p.numel() for p in m.parameters())}
98 return results
99
100if __name__=='__main__':
101 checks=math_checks()
102 try:
103 comparison=train_compare(); used_device=str(device)
104 except Exception as e:
105 print('CUDA/backend failure, falling back to CPU:', repr(e))
106 device=torch.device('cpu'); torch.manual_seed(SEED)
107 comparison=train_compare(); used_device='cpu (fallback)'
108 print(json.dumps({'device':used_device,'math_checks':checks,'comparison':comparison},indent=2))