Hypoelliptic transport-diffusion layer / experiment.py

Mechanism confirmed, baseline not beaten

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