cl33-oplm / so33.py
gman1911's picture
Selective reproducibility release for preprint v1.1 (One Object): two frozen 236M checkpoints, load-only model code, verified repro scripts, hashes
3ee235d verified
Raw History Blame Contribute Delete
2.69 kB
"""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<j — 6 compact rotations
(both axes same-sign) + 9 boosts (mixed-sign).
Reversibility (the interpretability handle): every rotor preserves η, so
R⁻¹ = η Rᵀ η exactly — no inverse needed. The scan is bit-reversible.
Everything here runs in float32 (bf16 diverges — proven on the lattice).
"""
from __future__ import annotations
import torch
ETA = torch.tensor([-1., -1., -1., 1., 1., 1.]) # (6,)
PAIRS = [(i, j) for i in range(6) for j in range(i + 1, 6)] # 15 pairs
N_GEN = len(PAIRS) # 15
def build_generators(device=None, dtype=torch.float32) -> 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"]