"""so(3,3) generator basis + matrix-exp rotor + reversible matrix-action scan. The operator-emitting LM works in a multi-block SO(3,3) state: each block is a 6-d vector; an emitted 15-coef bivector generates a per-block rotor R = exp(Ω), Ω ∈ so(3,3), and the state evolves by matrix action s_t = R_t · s_{t-1}. Signature η = diag(-1,-1,-1,+1,+1,+1) (matches cl33_t3lm/kge_inference.py). so(3,3) = { X : XᵀηX preserved } = { X : ηX antisymmetric }. A basis is G_ij = η (e_i e_jᵀ - e_j e_iᵀ) over the 15 pairs i torch.Tensor: """(15, 6, 6) so(3,3) generator matrices G_ij = η(e_i e_jᵀ - e_j e_iᵀ).""" eta = ETA.to(device=device, dtype=dtype) G = torch.zeros(N_GEN, 6, 6, device=device, dtype=dtype) for k, (i, j) in enumerate(PAIRS): A = torch.zeros(6, 6, device=device, dtype=dtype) A[i, j] = 1.0 A[j, i] = -1.0 G[k] = eta[:, None] * A # η @ A (η diagonal) return G def rotor_from_coefs(coefs: torch.Tensor, G: torch.Tensor) -> torch.Tensor: """coefs (..., 15) -> R (..., 6, 6) = matrix_exp(Σ coefs·G). Runs in the coefs' dtype; caller must keep this in float32. """ Omega = torch.einsum("...k,kij->...ij", coefs, G) # (...,6,6) in so(3,3) return torch.linalg.matrix_exp(Omega) def rotor_inverse(R: torch.Tensor) -> torch.Tensor: """R⁻¹ = η Rᵀ η — exact, no solve. R (...,6,6).""" eta = ETA.to(device=R.device, dtype=R.dtype) Rt = R.transpose(-1, -2) return eta[..., :, None] * Rt * eta[..., None, :] def eta_inner(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: """η-metric inner product ⟨a,b⟩_η over the last dim (6). a,b (...,6).""" eta = ETA.to(device=a.device, dtype=a.dtype) return (a * eta * b).sum(-1) def q_invariant(v: torch.Tensor) -> torch.Tensor: """Q(v) = -v0²-v1²-v2²+v3²+v4²+v5² (preserved by every rotor).""" return eta_inner(v, v) __all__ = ["ETA", "PAIRS", "N_GEN", "build_generators", "rotor_from_coefs", "rotor_inverse", "eta_inner", "q_invariant"]