"""Doubled-angle orientation order pooling.""" import torch def doubled_angle_pool(energies, eps=1e-8): """energies: (..., B), nonnegative; bins theta=b*pi/B. Returns normalized Cartesian order channels (zr, zi, q), plus raw z and mass. """ B = energies.shape[-1] theta = torch.arange(B, device=energies.device, dtype=energies.dtype) * torch.pi / B c, s = torch.cos(2 * theta), torch.sin(2 * theta) mass = energies.sum(-1) zr = (energies * c).sum(-1) zi = (energies * s).sum(-1) denom = mass + eps q = torch.sqrt(zr.square() + zi.square()) / denom return torch.stack((zr / denom, zi / denom, q), -1), (zr, zi, q, mass) def orientation_from_pool(pooled, q_min=1e-4): """Confidence-gated principal orientation in [0, pi).""" zr, zi, q = pooled[..., 0], pooled[..., 1], pooled[..., 2] phi = 0.5 * torch.atan2(zi, zr) % torch.pi return torch.where(q > q_min, phi, torch.zeros_like(phi))