"""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"]