Jacobian-Ranked Simplex Features / experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1import json, math, os, random
  2import numpy as np
  3import torch
  4from torch import nn
  5
  6SEED = 1729
  7
  8def seed_all(s):
  9    random.seed(s); np.random.seed(s); torch.manual_seed(s)
 10    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
 11
 12# Differentiable Heron area. Inputs are positive side lengths.
 13def heron(x):
 14    a,b,c = x[...,0], x[...,1], x[...,2]
 15    s = (a+b+c)/2
 16    q = (s*(s-a)*(s-b)*(s-c)).clamp_min(1e-12)
 17    return torch.sqrt(q)
 18
 19def triangle_data(n, seed):
 20    g = torch.Generator().manual_seed(seed)
 21    # Mixture includes ordinary and nearly collinear triangles, where ranking matters.
 22    x = torch.randn(n,3,3,generator=g)
 23    near = torch.rand(n,generator=g) < .25
 24    x[near,:,1] = x[near,:,0] + .03*torch.randn(near.sum(),3,generator=g)
 25    # Random per-example scale makes raw Jacobian sensitivity nonuniform.
 26    x *= (0.6 + 1.4*torch.rand(n,1,1,generator=g))
 27    p,q,r = x[:,0],x[:,1],x[:,2]
 28    sides = torch.stack([(p-q).norm(dim=1),(q-r).norm(dim=1),(r-p).norm(dim=1)],1)
 29    area = heron(sides)
 30    # sine of angle at p: an angle-sensitive, rigid-motion invariant target
 31    target = 2*area/(sides[:,0]*sides[:,2]).clamp_min(1e-8)
 32    target = target.clamp(0,1)
 33    return sides, area, target
 34
 35class MLP(nn.Module):
 36    def __init__(self,d):
 37        super().__init__(); self.net=nn.Sequential(nn.Linear(d,48),nn.Tanh(),nn.Linear(48,48),nn.Tanh(),nn.Linear(48,1))
 38    def forward(self,x): return self.net(x).squeeze(-1)
 39
 40def train_eval(kind, train, test, steps=500):
 41    ts,ta,ty=train; vs,va,vy=test
 42    def feat(s,a):
 43        if kind=='distance': return s
 44        if kind=='area': return torch.cat([s,a[:,None]],1)
 45        # Scalar triangle Jacobian has one singular value: norm of dA/d(a,b,c).
 46        z=s.detach().clone().requires_grad_(True)
 47        aa=heron(z); grad=torch.autograd.grad(aa.sum(),z)[0]
 48        sig=grad.norm(dim=1)
 49        w=(sig/(sig+1e-3)).clamp(0,1).detach()
 50        return torch.cat([s,(a*w)[:,None]],1)
 51    X,Y=feat(ts,ta),ty
 52    VX,VY=feat(vs,va),vy
 53    model=MLP(X.shape[1]); opt=torch.optim.Adam(model.parameters(),lr=3e-3)
 54    for _ in range(steps):
 55        opt.zero_grad(); loss=((model(X)-Y)**2).mean(); loss.backward(); opt.step()
 56    with torch.no_grad():
 57        mse=((model(VX)-VY)**2).mean().item()
 58        pred=model(VX)
 59    return mse, model, pred
 60
 61def matrix_check():
 62    # The paper's polynomial coordinate is alpha(x,y,z)=4*area^2 when
 63    # x,y,z are squared side lengths. Its gradient is the displayed row.
 64    t=torch.tensor([2.1,1.8,2.2,2.0,2.4,1.9,2.3,2.1,2.5],dtype=torch.double,requires_grad=True)
 65    tris=[(0,5,6),(1,2,8),(3,4,7),(6,7,8)]
 66    J=torch.zeros(4,9,dtype=torch.double)
 67    for row,ix in enumerate(tris):
 68        x=t[list(ix)]
 69        alpha=-0.5*(x*x).sum()+x[0]*x[1]+x[0]*x[2]+x[1]*x[2]
 70        J[row]=torch.autograd.grad(alpha,t,retain_graph=True)[0]
 71    M=torch.zeros(4,9,dtype=torch.double)
 72    M[0,[0,5,6]]=torch.stack([-t[0]+t[5]+t[6],t[0]-t[5]+t[6],t[0]+t[5]-t[6]])
 73    M[1,[1,2,8]]=torch.stack([-t[1]+t[2]+t[8],t[1]-t[2]+t[8],t[1]+t[2]-t[8]])
 74    M[2,[3,4,7]]=torch.stack([-t[3]+t[4]+t[7],t[3]-t[4]+t[7],t[3]+t[4]-t[7]])
 75    M[3,[6,7,8]]=torch.stack([-t[6]+t[7]+t[8],t[6]-t[7]+t[8],t[6]+t[7]-t[8]])
 76    err=(M-J).abs().max().item()
 77    sv=torch.linalg.svdvals(J).detach().numpy()
 78    return {'squared_length_polynomial_jacobian_max_abs_error':err,
 79            'displayed_matrix_singular_values':sv.tolist(),
 80            'displayed_matrix_rank':int((sv>1e-10).sum()),
 81            'smallest_singular_value':float(sv[-1])}
 82
 83def main():
 84    seed_all(SEED)
 85    check=matrix_check()
 86    results={k:[] for k in ['distance','area','weighted_area']}
 87    for split in range(5):
 88        train=triangle_data(700,11+2*split); test=triangle_data(500,12+2*split)
 89        for k in results:
 90            mse, model, pred=train_eval(k,train,test,steps=500)
 91            results[k].append(math.sqrt(mse))
 92    results={k:{'rmse_each_split':v,'mean_rmse':float(np.mean(v)),'std_rmse':float(np.std(v))} for k,v in results.items()}
 93    # Exact rigid transformation invariance of all geometric features.
 94    sides,area,target=test
 95    g=torch.Generator().manual_seed(99); R=torch.randn(3,3,generator=g); Q,_=torch.linalg.qr(R)
 96    # Reconstruct one set of points only for a numerical transformed-feature check.
 97    pts=torch.randn(200,3,3,generator=g); rot=pts@Q
 98    def ds(x):
 99        return torch.stack([(x[:,0]-x[:,1]).norm(dim=1),(x[:,1]-x[:,2]).norm(dim=1),(x[:,2]-x[:,0]).norm(dim=1)],1)
100    sdiff=(ds(pts)-ds(rot)).abs().max().item()
101    adiff=(heron(ds(pts))-heron(ds(rot))).abs().max().item()
102    out={'seed':SEED,'math_check':check,'benchmark':results,
103         'rigid_transform_max_distance_feature_change':sdiff,
104         'rigid_transform_max_area_feature_change':adiff,
105         'train_size':700,'test_size':500,'steps':500}
106    print(json.dumps(out,indent=2))
107
108if __name__=='__main__': main()