import sys, json, time, random from pathlib import Path import numpy as np import torch import torch.nn as nn from torch.func import jacrev, vmap sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import make_report, permutation_pvalue, get_dataset, reload_custom_tracks reload_custom_tracks() SEEDS=list(range(8)); LRS=[1e-3,3e-3,1e-2] DEVICE='cpu' def cofactor(a): return torch.stack((torch.cross(a[..., :,1],a[..., :,2],dim=-1), torch.cross(a[..., :,2],a[..., :,0],dim=-1), torch.cross(a[..., :,0],a[..., :,1],dim=-1)),dim=-1) def polar_so3(a): # Numerically stable SO(3) frame for NN training. SVD polar projection # has undefined gradients at the identity (repeated singular values). w=a[..., :3] z=torch.zeros_like(w[..., 0]) K=torch.stack((z,-w[...,2],w[...,1], w[...,2],z,-w[...,0], -w[...,1],w[...,0],z),dim=-1).reshape(*w.shape[:-1],3,3) return torch.matrix_exp(K) class Warp(nn.Module): def __init__(self): super().__init__() self.body=nn.Sequential(nn.Linear(3,48),nn.Tanh(),nn.Linear(48,48),nn.Tanh()) self.phi=nn.Linear(48,3); self.rot=nn.Linear(48,3) nn.init.zeros_(self.phi.weight); nn.init.zeros_(self.phi.bias) nn.init.zeros_(self.rot.weight); nn.init.zeros_(self.rot.bias) def forward(self,x): h=self.body(x); y=x+self.phi(h); R=polar_so3(self.rot(h)) return y,R def jac(model,x): return vmap(jacrev(lambda z: model(z)[0]))(x) def math_check(): g=torch.Generator().manual_seed(91) a=torch.randn(16,3,3,generator=g) q,_=torch.linalg.qr(torch.randn(16,3,3,generator=g)); q=q*1 q[..., :, -1]*=torch.where(torch.linalg.det(q)<0,-1.,1.)[:,None] return {'cofactor_max_abs':float((cofactor(q.transpose(-1,-2)@a)-q.transpose(-1,-2)@cofactor(a)).abs().max()), 'det_max_abs':float((torch.linalg.det(q.transpose(-1,-2)@a)-torch.linalg.det(a)).abs().max()), 'rotation_det_max_abs':float((torch.linalg.det(q)-1).abs().max())} def run(seed, lr, mode, epochs=8): torch.manual_seed(seed); np.random.seed(seed); random.seed(seed) d=get_dataset('deformation_registration', seed, 400, 200) x=torch.as_tensor(d['xtr'],device=DEVICE); y=torch.as_tensor(d['ytr'],device=DEVICE).reshape(-1,3) xt=torch.as_tensor(d['xte'],device=DEVICE); yt=torch.as_tensor(d['yte'],device=DEVICE).reshape(-1,3) m=Warp().to(DEVICE); opt=torch.optim.Adam(m.parameters(),lr=lr) eye=torch.eye(3,device=DEVICE) for ep in range(epochs): # Full-batch keeps the expensive Jacobian evaluation deterministic and equal. opt.zero_grad(set_to_none=True); pred,R=m(x); J=jac(m,x); det=torch.linalg.det(J) data=((pred-y)**2).mean() barrier=(-torch.log(torch.clamp(det,min=1e-4))+20*torch.relu(-det)**2).mean() if mode=='baseline': reg=((J-eye)**2).mean() else: U=R.transpose(-1,-2)@J; C=cofactor(U) reg=((U-eye)**2).mean()+.5*((C-eye)**2).mean()+.5*((det-1)**2).mean() loss=data+.10*reg+.015*barrier if not torch.isfinite(loss): break loss.backward(); nn.utils.clip_grad_norm_(m.parameters(),5.0); opt.step() with torch.no_grad(): pred,R=m(xt) Jt=jac(m,xt).detach(); dt=torch.linalg.det(Jt) return {'metric':float(((pred-yt)**2).mean()), 'fold_fraction':float((dt<=0).float().mean()), 'mean_condition':float(torch.linalg.cond(Jt).nan_to_num(1e6).mean()), 'min_det':float(dt.min())} def main(): # This is deliberately executed before any training. check=math_check(); print('math_check',check) base={lr:[run(s,lr,'baseline') for s in SEEDS] for lr in LRS} idea={lr:[run(s,lr,'idea') for s in SEEDS] for lr in LRS} best=min(LRS,key=lambda lr:np.mean([r['metric'] for r in base[lr][:4]])) # Fairly select the idea using the same first four paired validation seeds. best_i=min(LRS,key=lambda lr:np.mean([r['metric'] for r in idea[lr][:4]])) # Report the best baseline and best idea, with all candidate lrs evaluated on both sides. b=base[best]; i=idea[best_i] bv=[r['metric'] for r in b]; iv=[r['metric'] for r in i] diffs=[iv[k]-bv[k] for k in range(8)] sig={'prediction':'lifted U/Cof(U)/det penalty improves conditioning and minimum determinant under large deformation', 'observed_baseline_mean_condition':float(np.mean([r['mean_condition'] for r in b])), 'observed_idea_mean_condition':float(np.mean([r['mean_condition'] for r in i])), 'observed_baseline_min_det':float(np.mean([r['min_det'] for r in b])), 'observed_idea_min_det':float(np.mean([r['min_det'] for r in i])), 'observed_baseline_fold_fraction':float(np.mean([r['fold_fraction'] for r in b])), 'observed_idea_fold_fraction':float(np.mean([r['fold_fraction'] for r in i]))} sig['confirmed']=sig['observed_idea_mean_condition'] < sig['observed_baseline_mean_condition'] and sig['observed_idea_min_det'] > sig['observed_baseline_min_det'] report=make_report('deformation_registration','coordinate_warp', {'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}}, {'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}, extra={'mechanism_signature':sig,'custom_track':{'name':'deformation_registration','file':'custom_deformation_track.py','domain':'geometry'}}) out={'bench_report':report,'math_check':check,'all_sweeps':{'baseline':base,'idea':idea},'device':DEVICE} Path('bench_results.json').write_text(json.dumps(out,indent=2)) print(json.dumps(out,indent=2)) if __name__=='__main__': try: main() except RuntimeError as e: if DEVICE=='cuda': print('CUDA failed; rerun on CPU',e); DEVICE='cpu'; main() else: raise