Diffeomorphic gauge-fixing layer / experiment.py
Mechanism failed
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()