import json, math, random, time import numpy as np import torch import torch.nn as nn import torch.nn.functional as F SEED=483 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(4) device='cuda' if torch.cuda.is_available() else 'cpu' try: if device=='cuda': torch.cuda.get_device_properties(0) except Exception: device='cpu' # Spatial action for a translation g_t(x)=x+t: a(g_t,u)(x)=u(x-t). def grid_xy(n, dev): q=torch.linspace(-1,1,n,device=dev) y,x=torch.meshgrid(q,q,indexing='ij') return torch.stack((x,y),-1) def warp_translate(u,t): b,_,h,w=u.shape base=grid_xy(h,u.device).expand(b,-1,-1,-1) tn=torch.stack((2*t[:,0]/max(w-1,1),2*t[:,1]/max(h-1,1)),-1)[:,None,None,:] return F.grid_sample(u,base-tn,mode='bilinear',padding_mode='zeros',align_corners=True) def make_fields(n,b,dev): y,x=torch.meshgrid(torch.linspace(-1,1,n,device=dev),torch.linspace(-1,1,n,device=dev),indexing='ij') out=[] for _ in range(b): cx,cy=(torch.rand(2,device=dev)-.5)*.7 sx,sy=.16+torch.rand(2,device=dev)*.10 field=torch.exp(-((x-cx)**2/(2*sx*sx)+(y-cy)**2/(2*sy*sy))) out.append(field) return torch.stack(out)[:,None] def true_step(u): # deterministic translation law in canonical coordinates return warp_translate(u, torch.tensor([[1.3,-.8]],device=u.device).expand(u.shape[0],-1)) class ConvNet(nn.Module): def __init__(self): super().__init__() 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)) def forward(self,x): return self.net(x) class GaugeNet(nn.Module): def __init__(self,alpha=3.0): super().__init__(); self.alpha=alpha self.net=nn.Sequential(nn.Conv2d(1,8,5,padding=2),nn.ReLU(),nn.AdaptiveAvgPool2d(1),nn.Flatten(),nn.Linear(8,2)) def forward(self,u): return self.alpha*torch.tanh(self.net(u)) def center_offset(u): b,_,h,w=u.shape; yy,xx=torch.meshgrid(torch.arange(h,device=u.device),torch.arange(w,device=u.device),indexing='ij') mass=u.clamp_min(0)+1e-5; z=mass.sum((2,3)); return torch.stack(((mass*xx).sum((2,3))/z-w/2,(mass*yy).sum((2,3))/z-h/2),1) def train_baseline(x,y,steps=500): m=ConvNet().to(device); opt=torch.optim.Adam(m.parameters(),lr=2e-3) for _ in range(steps): p=m(x); loss=F.mse_loss(p,y); opt.zero_grad(); loss.backward(); opt.step() return m def train_gauge(x,y,ref,steps=500): g=GaugeNet().to(device); d=ConvNet().to(device); opt=torch.optim.Adam(list(g.parameters())+list(d.parameters()),lr=2e-3) for _ in range(steps): t=g(x); z=warp_translate(x,-t); pred=warp_translate(d(z),t) # Reference alignment is deliberately weak: it regularizes the estimated gauge. align=F.mse_loss(z,ref.expand_as(z)); loss=F.mse_loss(pred,y)+.03*align+.0005*(t*t).mean() opt.zero_grad(); loss.backward(); opt.step() return g,d def rollout_baseline(m,u,k): for _ in range(k): u=m(u) return u def rollout_gauge(g,d,u,k): for _ in range(k): t=g(u); u=warp_translate(d(warp_translate(u,-t)),t) return u def main(): global device n=24; train_b=64; test_b=128 # Probe the actual spatial operator and fall back if shared CUDA/cuDNN fails. if device == 'cuda': try: probe = torch.zeros(1,1,4,4,device='cuda') warp_translate(probe, torch.zeros(1,2,device='cuda')) except Exception: device = 'cpu' try: torch.cuda.empty_cache() except Exception: pass # Sanity check claimed group action and inverse composition for translations. 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) comp=(warp_translate(warp_translate(u,b),a)-warp_translate(u,a+b)).abs().mean().item() cyc=(warp_translate(warp_translate(u,a),-a)-u).abs().mean().item() # Data have independent coordinate gauges; target is the same transformed physical next state. base=make_fields(n,train_b,device); tr=(torch.rand(train_b,2,device=device)-.5)*5 x=warp_translate(base,tr); y=warp_translate(true_step(base),tr) test_base=make_fields(n,test_b,device); te=(torch.rand(test_b,2,device=device)-.5)*7 xt=warp_translate(test_base,te); yt=warp_translate(true_step(test_base),te) ref=base.mean(0,keepdim=True) t0=time.time(); bm=train_baseline(x,y); bt=time.time()-t0 t0=time.time(); gm,dm=train_gauge(x,y,ref); gt=time.time()-t0 with torch.no_grad(): bp= bm(xt); gp=rollout_gauge(gm,dm,xt.clone(),1) b20=rollout_baseline(bm,xt.clone(),20); g20=rollout_gauge(gm,dm,xt.clone(),20) # equivariance error under an unseen translation h h=torch.tensor([[2.2,-1.4]],device=device).expand(test_b,-1) eqb=(bm(warp_translate(xt,h))-warp_translate(bm(xt),h)).pow(2).mean().sqrt()/(bm(xt).pow(2).mean().sqrt()+1e-8) 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) 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())} print(json.dumps(metrics,indent=2)) if __name__=='__main__': main()