Polyconvex rotation-frame Jacobian loss / run_bench.py

Mechanism confirmed, baseline not beaten

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