cl33-oplm / wedge.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
2.95 kB
"""Exterior-product (wedge) tensors for the grade tower — metric-free.
The lattice grade tower is g3=⟨B·v⟩₃, g4=⟨B·B_prev⟩₄, g5=⟨B·B_prev·v⟩₅,
g6=⟨B·B_prev·B_prev2⟩₆. The top-grade part of a product of blades IS their
wedge, and the wedge (exterior product) is METRIC-FREE — pure antisymmetrization.
So the grade tower needs no Cayley/signature convention: fixed wedge tensors
built from blade combinatorics, computed in PARALLEL per position (not in the
sequential scan → no speed cost).
Blade ordering: grade-k blade = sorted k-subset of {0..5} in `combinations`
order. Grade-2 order matches so33.PAIRS (both are combinations(range(6),2)), so
the model's emitted 15 bivector coefs align with these tensors directly.
"""
from __future__ import annotations
from itertools import combinations
import torch
GRADE_DIM = {0: 1, 1: 6, 2: 15, 3: 20, 4: 15, 5: 6, 6: 1}
def blades(grade):
return list(combinations(range(6), grade))
def _shuffle_sign(A, B):
"""Sign of wedge of sorted disjoint tuples A,B = (-1)^#{(a,b): a>b}."""
if set(A) & set(B):
return 0, None
inv = sum(1 for a in A for b in B if a > b)
merged = tuple(sorted(A + B))
return (-1) ** inv, merged
def wedge_tensor(ga: int, gb: int) -> torch.Tensor:
"""(dim_{ga+gb}, dim_ga, dim_gb): C[c] = Σ W[c,a,b]·A[a]·B[b] realizes A∧B."""
A, B, C = blades(ga), blades(gb), blades(ga + gb)
Ci = {b: i for i, b in enumerate(C)}
W = torch.zeros(len(C), len(A), len(B))
for ia, a in enumerate(A):
for ib, b in enumerate(B):
s, m = _shuffle_sign(a, b)
if s != 0:
W[Ci[m], ia, ib] = float(s)
return W
class GradeTower(torch.nn.Module):
"""Given per-position grade-1 state v and emitted grade-2 operators B (with
causal shifts B_prev, B_prev2), produce the grade tower features.
Inputs (all (..., dim)): v (…,6), B (…,15) [B_prev/B_prev2 are B shifted]
Output: dict of g2..g6 features. g1=v and g0 handled by caller.
"""
def __init__(self):
super().__init__()
self.register_buffer("W_2_1", wedge_tensor(2, 1), persistent=False) # B∧v -> g3
self.register_buffer("W_2_2", wedge_tensor(2, 2), persistent=False) # B∧B -> g4
self.register_buffer("W_4_1", wedge_tensor(4, 1), persistent=False) # g4∧v -> g5
self.register_buffer("W_4_2", wedge_tensor(4, 2), persistent=False) # g4∧B2-> g6
def forward(self, v, B, B_prev, B_prev2):
g3 = torch.einsum("cab,...a,...b->...c", self.W_2_1, B, v) # (...,20)
g4 = torch.einsum("cab,...a,...b->...c", self.W_2_2, B, B_prev) # (...,15)
g5 = torch.einsum("cab,...a,...b->...c", self.W_4_1, g4, v) # (...,6)
g6 = torch.einsum("cab,...a,...b->...c", self.W_4_2, g4, B_prev2) # (...,1)
return g3, g4, g5, g6
__all__ = ["GRADE_DIM", "wedge_tensor", "GradeTower"]