Energy-trained monotone coordinate warp / experiment.py

✓✓ Beats tuned baseline

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