"""Exponential-map stochastic residual layer on S^2.""" import torch def project_tangent(x, u): return u - (u * x).sum(dim=-1, keepdim=True) * x def sphere_exp(x, v, eps=1e-12): # x is unit length and v tangent; stable closed-form exponential map. r = torch.linalg.vector_norm(v, dim=-1, keepdim=True) return torch.cos(r) * x + torch.sinc(r / torch.pi) * v def sphere_log(x, y, eps=1e-10): c = (x*y).sum(-1, keepdim=True).clamp(-1+eps, 1-eps) theta = torch.acos(c) w = y - c*x return theta * w / torch.linalg.vector_norm(w, dim=-1, keepdim=True).clamp_min(eps) def exp_residual(x, drift, h, sigma=0.0, noise=None): b = project_tangent(x, drift) if noise is None: noise = torch.randn_like(x) z = project_tangent(x, noise) v = h*b + (h**0.5)*sigma*z return sphere_exp(x, v) def stereographic(x): # chart from north pole, x3 != 1: q=(x1,x2)/(1-x3) return x[..., :2] / (1-x[..., 2:3]).clamp_min(1e-10) def inverse_stereographic(q): r2=(q*q).sum(-1, keepdim=True) return torch.cat((2*q/(1+r2), (r2-1)/(1+r2)), -1) def chart_jacobian(q): """Jacobian dx/dq, shape (...,3,2), for metric-aware chart update.""" a=q[...,0:1]; b=q[...,1:2]; d=1+(q*q).sum(-1,keepdim=True) # x=(2a/d,2b/d,(r2-1)/d) J=torch.zeros(q.shape[:-1]+(3,2),dtype=q.dtype,device=q.device) J[...,0,0]=2*(d[...,0]-2*a[...,0]*a[...,0])/d[...,0]**2 J[...,0,1]=-4*a[...,0]*b[...,0]/d[...,0]**2 J[...,1,0]=J[...,0,1]; J[...,1,1]=2*(d[...,0]-2*b[...,0]*b[...,0])/d[...,0]**2 J[...,2,0]=4*a[...,0]/d[...,0]**2; J[...,2,1]=4*b[...,0]/d[...,0]**2 return J