Polyconvex rotation-frame Jacobian loss / run_bench.py
Mechanism confirmed, baseline not beaten
1import sys, json, time, random
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6from torch.func import jacrev, vmap
7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
8from bench import make_report, permutation_pvalue, get_dataset, reload_custom_tracks
9reload_custom_tracks()
10
11SEEDS=list(range(8)); LRS=[1e-3,3e-3,1e-2]
12DEVICE='cpu'
13
14def cofactor(a):
15 return torch.stack((torch.cross(a[..., :,1],a[..., :,2],dim=-1),
16 torch.cross(a[..., :,2],a[..., :,0],dim=-1),
17 torch.cross(a[..., :,0],a[..., :,1],dim=-1)),dim=-1)
18
19def polar_so3(a):
20 # Numerically stable SO(3) frame for NN training. SVD polar projection
21 # has undefined gradients at the identity (repeated singular values).
22 w=a[..., :3]
23 z=torch.zeros_like(w[..., 0])
24 K=torch.stack((z,-w[...,2],w[...,1], w[...,2],z,-w[...,0],
25 -w[...,1],w[...,0],z),dim=-1).reshape(*w.shape[:-1],3,3)
26 return torch.matrix_exp(K)
27
28class Warp(nn.Module):
29 def __init__(self):
30 super().__init__()
31 self.body=nn.Sequential(nn.Linear(3,48),nn.Tanh(),nn.Linear(48,48),nn.Tanh())
32 self.phi=nn.Linear(48,3); self.rot=nn.Linear(48,3)
33 nn.init.zeros_(self.phi.weight); nn.init.zeros_(self.phi.bias)
34 nn.init.zeros_(self.rot.weight); nn.init.zeros_(self.rot.bias)
35 def forward(self,x):
36 h=self.body(x); y=x+self.phi(h); R=polar_so3(self.rot(h))
37 return y,R
38
39def jac(model,x):
40 return vmap(jacrev(lambda z: model(z)[0]))(x)
41
42def math_check():
43 g=torch.Generator().manual_seed(91)
44 a=torch.randn(16,3,3,generator=g)
45 q,_=torch.linalg.qr(torch.randn(16,3,3,generator=g)); q=q*1
46 q[..., :, -1]*=torch.where(torch.linalg.det(q)<0,-1.,1.)[:,None]
47 return {'cofactor_max_abs':float((cofactor(q.transpose(-1,-2)@a)-q.transpose(-1,-2)@cofactor(a)).abs().max()),
48 'det_max_abs':float((torch.linalg.det(q.transpose(-1,-2)@a)-torch.linalg.det(a)).abs().max()),
49 'rotation_det_max_abs':float((torch.linalg.det(q)-1).abs().max())}
50
51def run(seed, lr, mode, epochs=8):
52 torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
53 d=get_dataset('deformation_registration', seed, 400, 200)
54 x=torch.as_tensor(d['xtr'],device=DEVICE); y=torch.as_tensor(d['ytr'],device=DEVICE).reshape(-1,3)
55 xt=torch.as_tensor(d['xte'],device=DEVICE); yt=torch.as_tensor(d['yte'],device=DEVICE).reshape(-1,3)
56 m=Warp().to(DEVICE); opt=torch.optim.Adam(m.parameters(),lr=lr)
57 eye=torch.eye(3,device=DEVICE)
58 for ep in range(epochs):
59 # Full-batch keeps the expensive Jacobian evaluation deterministic and equal.
60 opt.zero_grad(set_to_none=True); pred,R=m(x); J=jac(m,x); det=torch.linalg.det(J)
61 data=((pred-y)**2).mean()
62 barrier=(-torch.log(torch.clamp(det,min=1e-4))+20*torch.relu(-det)**2).mean()
63 if mode=='baseline': reg=((J-eye)**2).mean()
64 else:
65 U=R.transpose(-1,-2)@J; C=cofactor(U)
66 reg=((U-eye)**2).mean()+.5*((C-eye)**2).mean()+.5*((det-1)**2).mean()
67 loss=data+.10*reg+.015*barrier
68 if not torch.isfinite(loss): break
69 loss.backward(); nn.utils.clip_grad_norm_(m.parameters(),5.0); opt.step()
70 with torch.no_grad(): pred,R=m(xt)
71 Jt=jac(m,xt).detach(); dt=torch.linalg.det(Jt)
72 return {'metric':float(((pred-yt)**2).mean()),
73 'fold_fraction':float((dt<=0).float().mean()),
74 'mean_condition':float(torch.linalg.cond(Jt).nan_to_num(1e6).mean()),
75 'min_det':float(dt.min())}
76
77def main():
78 # This is deliberately executed before any training.
79 check=math_check(); print('math_check',check)
80 base={lr:[run(s,lr,'baseline') for s in SEEDS] for lr in LRS}
81 idea={lr:[run(s,lr,'idea') for s in SEEDS] for lr in LRS}
82 best=min(LRS,key=lambda lr:np.mean([r['metric'] for r in base[lr][:4]]))
83 # Fairly select the idea using the same first four paired validation seeds.
84 best_i=min(LRS,key=lambda lr:np.mean([r['metric'] for r in idea[lr][:4]]))
85 # Report the best baseline and best idea, with all candidate lrs evaluated on both sides.
86 b=base[best]; i=idea[best_i]
87 bv=[r['metric'] for r in b]; iv=[r['metric'] for r in i]
88 diffs=[iv[k]-bv[k] for k in range(8)]
89 sig={'prediction':'lifted U/Cof(U)/det penalty improves conditioning and minimum determinant under large deformation',
90 'observed_baseline_mean_condition':float(np.mean([r['mean_condition'] for r in b])),
91 'observed_idea_mean_condition':float(np.mean([r['mean_condition'] for r in i])),
92 'observed_baseline_min_det':float(np.mean([r['min_det'] for r in b])),
93 'observed_idea_min_det':float(np.mean([r['min_det'] for r in i])),
94 'observed_baseline_fold_fraction':float(np.mean([r['fold_fraction'] for r in b])),
95 'observed_idea_fold_fraction':float(np.mean([r['fold_fraction'] for r in i]))}
96 sig['confirmed']=sig['observed_idea_mean_condition'] < sig['observed_baseline_mean_condition'] and sig['observed_idea_min_det'] > sig['observed_baseline_min_det']
97 report=make_report('deformation_registration','coordinate_warp',
98 {'best_config':{'lr':best,'epochs':8},'sweep':{str(k):[r['metric'] for r in v] for k,v in base.items()},'full':{'per_seed':bv}},
99 {'best_config':{'lr':best_i,'epochs':8},'sweep':{str(k):[r['metric'] for r in v] for k,v in idea.items()},'per_seed':iv},
100 extra={'mechanism_signature':sig,'custom_track':{'name':'deformation_registration','file':'custom_deformation_track.py','domain':'geometry'}})
101 out={'bench_report':report,'math_check':check,'all_sweeps':{'baseline':base,'idea':idea},'device':DEVICE}
102 Path('bench_results.json').write_text(json.dumps(out,indent=2))
103 print(json.dumps(out,indent=2))
104
105if __name__=='__main__':
106 try: main()
107 except RuntimeError as e:
108 if DEVICE=='cuda':
109 print('CUDA failed; rerun on CPU',e); DEVICE='cpu'; main()
110 else: raise