Jacobian-Ranked Simplex Features / experiment.py
Beats tuned baseline
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()