Exponential-Map Stochastic Residual Layer / sphere_exp_residual.py
Beats tuned baseline
1"""Exponential-map stochastic residual layer on S^2."""
2import torch
3
4
5def project_tangent(x, u):
6 return u - (u * x).sum(dim=-1, keepdim=True) * x
7
8
9def sphere_exp(x, v, eps=1e-12):
10 # x is unit length and v tangent; stable closed-form exponential map.
11 r = torch.linalg.vector_norm(v, dim=-1, keepdim=True)
12 return torch.cos(r) * x + torch.sinc(r / torch.pi) * v
13
14
15def sphere_log(x, y, eps=1e-10):
16 c = (x*y).sum(-1, keepdim=True).clamp(-1+eps, 1-eps)
17 theta = torch.acos(c)
18 w = y - c*x
19 return theta * w / torch.linalg.vector_norm(w, dim=-1, keepdim=True).clamp_min(eps)
20
21
22def exp_residual(x, drift, h, sigma=0.0, noise=None):
23 b = project_tangent(x, drift)
24 if noise is None:
25 noise = torch.randn_like(x)
26 z = project_tangent(x, noise)
27 v = h*b + (h**0.5)*sigma*z
28 return sphere_exp(x, v)
29
30
31def stereographic(x):
32 # chart from north pole, x3 != 1: q=(x1,x2)/(1-x3)
33 return x[..., :2] / (1-x[..., 2:3]).clamp_min(1e-10)
34
35
36def inverse_stereographic(q):
37 r2=(q*q).sum(-1, keepdim=True)
38 return torch.cat((2*q/(1+r2), (r2-1)/(1+r2)), -1)
39
40
41def chart_jacobian(q):
42 """Jacobian dx/dq, shape (...,3,2), for metric-aware chart update."""
43 a=q[...,0:1]; b=q[...,1:2]; d=1+(q*q).sum(-1,keepdim=True)
44 # x=(2a/d,2b/d,(r2-1)/d)
45 J=torch.zeros(q.shape[:-1]+(3,2),dtype=q.dtype,device=q.device)
46 J[...,0,0]=2*(d[...,0]-2*a[...,0]*a[...,0])/d[...,0]**2
47 J[...,0,1]=-4*a[...,0]*b[...,0]/d[...,0]**2
48 J[...,1,0]=J[...,0,1]; J[...,1,1]=2*(d[...,0]-2*b[...,0]*b[...,0])/d[...,0]**2
49 J[...,2,0]=4*a[...,0]/d[...,0]**2; J[...,2,1]=4*b[...,0]/d[...,0]**2
50 return J