staging / src /optimizer /manifolds.py
kinaar111's picture
Staging: mystery shell + 6-string optimizer + physics-audio
97bec8a verified
Raw
History Blame Contribute Delete
3.12 kB
# src/optimizer/manifolds.py
"""Manifold helpers: stereographic S^3 projection and Fisher-Rao simplex geometry."""
from __future__ import annotations
import sys
from pathlib import Path
import torch
import torch.nn.functional as F
import geoopt
def _ensure_fisher_rao_path() -> None:
candidates = [
Path(__file__).resolve().parents[2] / "fisher_rao",
Path.home() / "Projects" / "Fisher_Rao",
]
for root in candidates:
if root.exists() and str(root.parent) not in sys.path:
sys.path.insert(0, str(root.parent))
return
_ensure_fisher_rao_path()
try:
from fisher_rao.metrics import (
sqrt_embed,
simplex_to_sphere,
sphere_to_simplex,
softmax_to_fr_sphere,
fisher_rao_distance,
egrad_to_rgrad_fr,
)
except ImportError:
def _normalize_probs(p: torch.Tensor, eps: float = 1e-8) -> torch.Tensor:
p = p.clamp(min=eps)
return p / p.sum(dim=-1, keepdim=True)
def sqrt_embed(p: torch.Tensor, eps: float = 1e-8) -> torch.Tensor:
p = _normalize_probs(p, eps=eps)
s = p.sqrt()
return s / s.norm(dim=-1, keepdim=True).clamp(min=eps)
simplex_to_sphere = sqrt_embed
def sphere_to_simplex(s: torch.Tensor, eps: float = 1e-8) -> torch.Tensor:
return _normalize_probs(s.pow(2).clamp(min=0.0), eps=eps)
def softmax_to_fr_sphere(logits: torch.Tensor) -> torch.Tensor:
return sqrt_embed(F.softmax(logits, dim=-1))
def fisher_rao_distance(p: torch.Tensor, q: torch.Tensor, eps: float = 1e-8) -> torch.Tensor:
p = _normalize_probs(p, eps=eps)
q = _normalize_probs(q, eps=eps)
bc = (p.sqrt() * q.sqrt()).sum(dim=-1).clamp(0.0, 1.0)
return 2.0 * torch.acos(bc.clamp(-1.0 + 1e-7, 1.0 - 1e-7))
def egrad_to_rgrad_fr(s: torch.Tensor, egrad: torch.Tensor) -> torch.Tensor:
inner = (s * egrad).sum(dim=-1, keepdim=True)
return egrad - inner * s
class FisherRaoSphere(geoopt.manifolds.Sphere):
"""
Unit sphere with Fisher-Rao interpretation via sqrt-embedding of simplices.
Geodesics and inner products are the standard round metric on S^{n-1},
which corresponds to the Fisher-Rao metric on Δ^{n-1} under φ(p) = √p/||√p||.
"""
def simplex_to_manifold(self, p: torch.Tensor) -> torch.Tensor:
return simplex_to_sphere(p)
def manifold_to_simplex(self, s: torch.Tensor) -> torch.Tensor:
return sphere_to_simplex(s)
def fr_distance_on_simplex(self, p: torch.Tensor, q: torch.Tensor) -> torch.Tensor:
return fisher_rao_distance(p, q)
def egrad2rgrad(self, x: torch.Tensor, u: torch.Tensor) -> torch.Tensor:
return egrad_to_rgrad_fr(x, u)
def logits_to_fr_manifold_param(
logits: torch.Tensor,
manifold: FisherRaoSphere | None = None,
) -> geoopt.ManifoldParameter:
"""Create a ManifoldParameter on the Fisher-Rao sphere from unconstrained logits."""
if manifold is None:
manifold = FisherRaoSphere()
s = softmax_to_fr_sphere(logits)
return geoopt.ManifoldParameter(s, manifold=manifold)