Diffeomorphic gauge-fixing layer / experiment.py

Mechanism failed

Raw ⬇ ZIP
  1import json, math, random, time
  2import numpy as np
  3import torch
  4import torch.nn as nn
  5import torch.nn.functional as F
  6
  7SEED=483
  8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
  9torch.set_num_threads(4)
 10device='cuda' if torch.cuda.is_available() else 'cpu'
 11try:
 12    if device=='cuda': torch.cuda.get_device_properties(0)
 13except Exception:
 14    device='cpu'
 15
 16# Spatial action for a translation g_t(x)=x+t: a(g_t,u)(x)=u(x-t).
 17def grid_xy(n, dev):
 18    q=torch.linspace(-1,1,n,device=dev)
 19    y,x=torch.meshgrid(q,q,indexing='ij')
 20    return torch.stack((x,y),-1)
 21
 22def warp_translate(u,t):
 23    b,_,h,w=u.shape
 24    base=grid_xy(h,u.device).expand(b,-1,-1,-1)
 25    tn=torch.stack((2*t[:,0]/max(w-1,1),2*t[:,1]/max(h-1,1)),-1)[:,None,None,:]
 26    return F.grid_sample(u,base-tn,mode='bilinear',padding_mode='zeros',align_corners=True)
 27
 28def make_fields(n,b,dev):
 29    y,x=torch.meshgrid(torch.linspace(-1,1,n,device=dev),torch.linspace(-1,1,n,device=dev),indexing='ij')
 30    out=[]
 31    for _ in range(b):
 32        cx,cy=(torch.rand(2,device=dev)-.5)*.7
 33        sx,sy=.16+torch.rand(2,device=dev)*.10
 34        field=torch.exp(-((x-cx)**2/(2*sx*sx)+(y-cy)**2/(2*sy*sy)))
 35        out.append(field)
 36    return torch.stack(out)[:,None]
 37
 38def true_step(u):
 39    # deterministic translation law in canonical coordinates
 40    return warp_translate(u, torch.tensor([[1.3,-.8]],device=u.device).expand(u.shape[0],-1))
 41
 42class ConvNet(nn.Module):
 43    def __init__(self):
 44        super().__init__()
 45        self.net=nn.Sequential(nn.Conv2d(1,16,5,padding=2),nn.ReLU(),nn.Conv2d(16,16,3,padding=1),nn.ReLU(),nn.Conv2d(16,1,3,padding=1))
 46    def forward(self,x): return self.net(x)
 47
 48class GaugeNet(nn.Module):
 49    def __init__(self,alpha=3.0):
 50        super().__init__(); self.alpha=alpha
 51        self.net=nn.Sequential(nn.Conv2d(1,8,5,padding=2),nn.ReLU(),nn.AdaptiveAvgPool2d(1),nn.Flatten(),nn.Linear(8,2))
 52    def forward(self,u): return self.alpha*torch.tanh(self.net(u))
 53
 54def center_offset(u):
 55    b,_,h,w=u.shape; yy,xx=torch.meshgrid(torch.arange(h,device=u.device),torch.arange(w,device=u.device),indexing='ij')
 56    mass=u.clamp_min(0)+1e-5; z=mass.sum((2,3));
 57    return torch.stack(((mass*xx).sum((2,3))/z-w/2,(mass*yy).sum((2,3))/z-h/2),1)
 58
 59def train_baseline(x,y,steps=500):
 60    m=ConvNet().to(device); opt=torch.optim.Adam(m.parameters(),lr=2e-3)
 61    for _ in range(steps):
 62        p=m(x); loss=F.mse_loss(p,y); opt.zero_grad(); loss.backward(); opt.step()
 63    return m
 64
 65def train_gauge(x,y,ref,steps=500):
 66    g=GaugeNet().to(device); d=ConvNet().to(device); opt=torch.optim.Adam(list(g.parameters())+list(d.parameters()),lr=2e-3)
 67    for _ in range(steps):
 68        t=g(x); z=warp_translate(x,-t); pred=warp_translate(d(z),t)
 69        # Reference alignment is deliberately weak: it regularizes the estimated gauge.
 70        align=F.mse_loss(z,ref.expand_as(z)); loss=F.mse_loss(pred,y)+.03*align+.0005*(t*t).mean()
 71        opt.zero_grad(); loss.backward(); opt.step()
 72    return g,d
 73
 74def rollout_baseline(m,u,k):
 75    for _ in range(k): u=m(u)
 76    return u
 77
 78def rollout_gauge(g,d,u,k):
 79    for _ in range(k):
 80        t=g(u); u=warp_translate(d(warp_translate(u,-t)),t)
 81    return u
 82
 83def main():
 84    global device
 85    n=24; train_b=64; test_b=128
 86    # Probe the actual spatial operator and fall back if shared CUDA/cuDNN fails.
 87    if device == 'cuda':
 88        try:
 89            probe = torch.zeros(1,1,4,4,device='cuda')
 90            warp_translate(probe, torch.zeros(1,2,device='cuda'))
 91        except Exception:
 92            device = 'cpu'
 93            try:
 94                torch.cuda.empty_cache()
 95            except Exception:
 96                pass
 97    # Sanity check claimed group action and inverse composition for translations.
 98    u=make_fields(n,16,device); a=torch.tensor([[1.7,-1.1]],device=device).expand(16,-1); b=torch.tensor([[-.8,.6]],device=device).expand(16,-1)
 99    comp=(warp_translate(warp_translate(u,b),a)-warp_translate(u,a+b)).abs().mean().item()
100    cyc=(warp_translate(warp_translate(u,a),-a)-u).abs().mean().item()
101    # Data have independent coordinate gauges; target is the same transformed physical next state.
102    base=make_fields(n,train_b,device); tr=(torch.rand(train_b,2,device=device)-.5)*5
103    x=warp_translate(base,tr); y=warp_translate(true_step(base),tr)
104    test_base=make_fields(n,test_b,device); te=(torch.rand(test_b,2,device=device)-.5)*7
105    xt=warp_translate(test_base,te); yt=warp_translate(true_step(test_base),te)
106    ref=base.mean(0,keepdim=True)
107    t0=time.time(); bm=train_baseline(x,y); bt=time.time()-t0
108    t0=time.time(); gm,dm=train_gauge(x,y,ref); gt=time.time()-t0
109    with torch.no_grad():
110        bp= bm(xt); gp=rollout_gauge(gm,dm,xt.clone(),1)
111        b20=rollout_baseline(bm,xt.clone(),20); g20=rollout_gauge(gm,dm,xt.clone(),20)
112        # equivariance error under an unseen translation h
113        h=torch.tensor([[2.2,-1.4]],device=device).expand(test_b,-1)
114        eqb=(bm(warp_translate(xt,h))-warp_translate(bm(xt),h)).pow(2).mean().sqrt()/(bm(xt).pow(2).mean().sqrt()+1e-8)
115        eqg=(rollout_gauge(gm,dm,warp_translate(xt,h),1)-warp_translate(gp,h)).pow(2).mean().sqrt()/(gp.pow(2).mean().sqrt()+1e-8)
116        metrics={'baseline_one_step_mse':F.mse_loss(bp,yt).item(),'gauge_one_step_mse':F.mse_loss(gp,yt).item(),'baseline_20_step_mse':F.mse_loss(b20,yt).item(),'gauge_20_step_mse':F.mse_loss(g20,yt).item(),'baseline_equivariance':eqb.item(),'gauge_equivariance':eqg.item(),'group_action_composition_abs':comp,'inverse_cycle_abs':cyc,'baseline_seconds':bt,'gauge_seconds':gt,'device':device,'params_baseline':sum(p.numel() for p in bm.parameters()),'params_gauge_total':sum(p.numel() for p in gm.parameters())+sum(p.numel() for p in dm.parameters())}
117    print(json.dumps(metrics,indent=2))
118
119if __name__=='__main__': main()