Energy-trained monotone coordinate warp / experiment.py
Beats tuned baseline
1import json, math, time
2import numpy as np
3import torch
4from torch import nn
5
6SEED=1310
7np.random.seed(SEED); torch.manual_seed(SEED)
8DT=np.float64
9
10def warp_numpy(q, n=20001, eps=1e-12, h=None):
11 s=np.linspace(0,1,n)
12 if h is None: h=np.zeros_like(s)
13 rho=eps + np.maximum(s,1e-15)**(q-1)*np.exp(np.clip(h,-5,5))
14 # trapezoidal cumulative integration
15 R=np.concatenate([[0.], np.cumsum((rho[:-1]+rho[1:])*.5*(s[1]-s[0]))])
16 r=R/R[-1]
17 return s,r,rho/R[-1]
18
19def log_slope(x,y,lo=1e-3,hi=.1):
20 m=(x>lo)&(x<hi)&(y>0)
21 return float(np.polyfit(np.log(x[m]),np.log(y[m]),1)[0])
22
23def interp_max_error(q, lam, n):
24 sg,rg,_=warp_numpy(q,n=50001)
25 # target after mapping, sampled on computational grid
26 y=rg**lam
27 # evaluate exact on dense points and linear interpolation from n nodes
28 sn=np.linspace(0,1,n); rn=np.interp(sn,sg,rg); yn=rn**lam
29 xx=np.linspace(0,1,100001)
30 yy=np.interp(xx,sn,yn)
31 exact=np.interp(xx,sg,y)
32 return float(np.max(np.abs(yy-exact)))
33
34class MLP(nn.Module):
35 def __init__(self):
36 super().__init__(); self.net=nn.Sequential(nn.Linear(1,32),nn.Tanh(),nn.Linear(32,32),nn.Tanh(),nn.Linear(32,1))
37 def forward(self,x): return self.net(x)
38
39class LearnedWarp(nn.Module):
40 def __init__(self,q=2,ng=257):
41 super().__init__(); self.q=q; self.grid=torch.linspace(0,1,ng).view(-1,1)
42 self.h=nn.Sequential(nn.Linear(1,12),nn.Tanh(),nn.Linear(12,1))
43 def map(self,s):
44 # differentiable cumulative trapezoids on fixed grid, then linear interpolation
45 g=self.grid.to(s.device); hh=torch.clamp(self.h(g),-2.0,2.0)
46 rho=1e-5 + torch.clamp(g,min=1e-6)**(self.q-1)*torch.exp(hh)
47 d=g[1:]-g[:-1]; inc=.5*(rho[1:]+rho[:-1])*d
48 R=torch.cat([torch.zeros(1,1,device=s.device),torch.cumsum(inc,0)],0); R=R/R[-1]
49 idx=torch.clamp((s*(len(g)-1)).long(),0,len(g)-2); t=s*(len(g)-1)-idx
50 return R[idx,0]*(1-t)+R[idx+1,0]*t
51
52def train_case(kind, steps=500):
53 torch.manual_seed(SEED+({'identity':0,'fixed':1,'learned':2}[kind]))
54 model=MLP(); params=list(model.parameters())
55 warp=LearnedWarp() if kind=='learned' else None
56 if warp: params += list(warp.parameters())
57 opt=torch.optim.Adam(params,lr=2e-3)
58 # fixed samples, held-out deterministic grid
59 s=torch.linspace(.0005,1,96).view(-1,1); target_phys=s**.5
60 sh=torch.linspace(.002,1,400).view(-1,1); yh=sh**.5
61 t0=time.time()
62 for k in range(steps):
63 if kind=='identity': x=s; y=target_phys
64 elif kind=='fixed': x=s; y=s # r=s^2, physical sqrt(r)=s
65 else:
66 r=warp.map(s); x=s; y=torch.sqrt(r)
67 pred=model(x); loss=((pred-y)**2).mean()
68 # weak regularity control on learned h sampled grid
69 if warp: loss=loss+1e-5*(warp.h(warp.grid[1:])-warp.h(warp.grid[:-1])).pow(2).mean()
70 opt.zero_grad(); loss.backward(); opt.step()
71 with torch.no_grad():
72 if kind=='identity': yt=yh
73 elif kind=='fixed': yt=sh
74 else: yt=torch.sqrt(warp.map(sh))
75 err=((model(sh)-yt)**2).mean().sqrt().item()
76 return err, time.time()-t0
77
78def main():
79 # Prediction 1: exact endpoint normalization, strict positivity, and derivative identity.
80 endpoint=[]
81 for q in [1,1.5,2,3,4]:
82 s,r,rp=warp_numpy(q)
83 fd=np.gradient(r,s); rel=float(np.max(np.abs(fd[2:-2]-rp[2:-2])/(np.abs(rp[2:-2])+1e-12)))
84 endpoint.append({'q':q,'r0':float(r[0]),'r1':float(r[-1]),'min_rprime':float(rp[1:].min()),'fitted_r_exponent':log_slope(s,r),'max_relative_derivative_error':rel})
85 # Prediction 2: exponent multiplication over q,lambda sweep.
86 exponent=[]
87 for q in [1,1.5,2,3,4]:
88 for lam in [.25,.5,.75]:
89 s,r,_=warp_numpy(q); exponent.append({'q':q,'lambda':lam,'predicted':q*lam,'observed':log_slope(s,r**lam)})
90 # Prediction 3: higher transformed regularity gives steeper interpolation convergence.
91 conv={}
92 for q in [1,2,3]:
93 lam=.25; ns=[17,33,65,129,257]
94 es=[interp_max_error(q,lam,n) for n in ns]
95 slope=float(np.polyfit(np.log(ns),np.log(es),1)[0])
96 conv[str(q)]={'q_lambda':q*lam,'ns':ns,'errors':es,'observed_error_slope_vs_N':slope,'predicted_slope_ideal':-min(2,q*lam)}
97 mlp={}
98 for k in ['identity','fixed','learned']:
99 e,t=train_case(k); mlp[k]={'heldout_rmse':e,'seconds':t}
100 out={'seed':SEED,'predictions':{'monotonic_normalized':endpoint,'exponent_multiplication':exponent,'interpolation_convergence':conv},'mini_experiment':mlp}
101 with open('results.json','w') as f: json.dump(out,f,indent=2)
102 print(json.dumps(out,indent=2))
103
104if __name__=='__main__': main()