"""Operator-emitting language model (cl33-opLM v0). The OPERATOR-ONLY thesis, made structural: context reaches the next-token prediction ONLY through emitted operators. A causal transformer "emitter" produces, per position, a per-block so(3,3) bivector; those generate rotors; a reversible matrix-action scan evolves a multi-block SO(3,3) state; the readout sees ONLY that state. No residual bypass from the emitter to the readout — so if the operators don't carry the information, PPL suffers (that IS the thesis). tokens → embed → causal emitter → h_t h_t → op_head → coefs_t (n_blocks, 15) [clipped, small init] R_t = matrix_exp(Σ coefs_t·G) [per block, fp32] s_t = R_t · s_{t-1} [reversible scan; s_0 learned] logits = readout(LayerNorm(flatten s_t)) [readout sees ONLY the state] Scan runs in float32 (bf16 diverges — proven on the lattice). Emitter may be autocast-bf16; the state path is force-fp32. """ from __future__ import annotations import math import sys from dataclasses import dataclass from pathlib import Path import torch import torch.nn as nn import torch.nn.functional as F sys.path.insert(0, str(Path(__file__).resolve().parent)) from so33 import build_generators, rotor_from_coefs, N_GEN # type: ignore from wedge import GradeTower # type: ignore # per-block readout feature dims: g1 + g2 + g3 + g4 + g5 + g6 _GRADE1, _GRADE2, _GRADE3, _GRADE4, _GRADE5, _GRADE6 = 6, 15, 20, 15, 6, 1 _TOWER_DIM = _GRADE1 + _GRADE2 + _GRADE3 + _GRADE4 + _GRADE5 + _GRADE6 # 63 @dataclass class OpEmitConfig: vocab_size: int = 8192 d_model: int = 384 n_layers: int = 6 n_heads: int = 6 d_ff: int = 1536 max_seq_len: int = 256 n_blocks: int = 16 # multi-block SO(3,3) state (6·n_blocks dims) coef_clip: float = 1.0 # per-block bivector L2-norm cap (bounds growth) op_init_scale: float = 0.02 # rotors start ≈ identity use_grade_tower: bool = True # readout sees full grade tower per block dropout: float = 0.0 class EmitterBlock(nn.Module): def __init__(self, c: OpEmitConfig): super().__init__() self.n_heads = c.n_heads self.d_head = c.d_model // c.n_heads self.qkv = nn.Linear(c.d_model, 3 * c.d_model, bias=False) self.o = nn.Linear(c.d_model, c.d_model, bias=False) self.norm1 = nn.LayerNorm(c.d_model) self.norm2 = nn.LayerNorm(c.d_model) self.mlp = nn.Sequential(nn.Linear(c.d_model, c.d_ff), nn.GELU(), nn.Linear(c.d_ff, c.d_model)) def forward(self, x): B, T, D = x.shape h = self.norm1(x) qkv = self.qkv(h).reshape(B, T, 3, self.n_heads, self.d_head) q, k, v = qkv.unbind(2) q, k, v = (t.transpose(1, 2) for t in (q, k, v)) # (B,H,T,dh) a = F.scaled_dot_product_attention(q, k, v, is_causal=True) a = a.transpose(1, 2).reshape(B, T, D) x = x + self.o(a) x = x + self.mlp(self.norm2(x)) return x class OpEmitLM(nn.Module): def __init__(self, c: OpEmitConfig): super().__init__() self.c = c self.tok_embed = nn.Embedding(c.vocab_size, c.d_model) self.pos_embed = nn.Embedding(c.max_seq_len, c.d_model) self.blocks = nn.ModuleList([EmitterBlock(c) for _ in range(c.n_layers)]) self.norm = nn.LayerNorm(c.d_model) self.op_head = nn.Linear(c.d_model, c.n_blocks * N_GEN) nn.init.normal_(self.op_head.weight, std=1e-3) nn.init.zeros_(self.op_head.bias) # learned initial state per block (nonzero so rotors have something to act on) self.s0 = nn.Parameter(torch.randn(c.n_blocks, 6) * 0.5) feat_per_block = _TOWER_DIM if c.use_grade_tower else 6 self.feat_dim = c.n_blocks * feat_per_block self.tower = GradeTower() if c.use_grade_tower else None self.state_norm = nn.LayerNorm(self.feat_dim) self.readout = nn.Linear(self.feat_dim, c.vocab_size, bias=False) self.register_buffer("G", build_generators(dtype=torch.float32), persistent=False) def emit(self, idx): B, T = idx.shape pos = torch.arange(T, device=idx.device) x = self.tok_embed(idx) + self.pos_embed(pos)[None] for blk in self.blocks: x = blk(x) h = self.norm(x) coefs = self.op_head(h).view(B, T, self.c.n_blocks, N_GEN) coefs = coefs * self.c.op_init_scale # per-block L2-norm clip → bounded rotors → bounded state growth n = coefs.norm(dim=-1, keepdim=True) coefs = coefs * (self.c.coef_clip / n.clamp_min(self.c.coef_clip)) return coefs # (B,T,nb,15) def scan(self, coefs, ablate=False): """Reversible matrix-action scan. Returns S (B,T,nb,6). fp32 forced.""" B, T, nb, _ = coefs.shape coefs = coefs.float() if ablate: R = torch.eye(6, device=coefs.device).expand(B, T, nb, 6, 6) else: R = rotor_from_coefs(coefs, self.G.float()) # (B,T,nb,6,6) s = self.s0.float().expand(B, nb, 6).contiguous() s = s / s.norm(dim=-1, keepdim=True).clamp_min(1e-6) S = torch.empty(B, T, nb, 6, device=coefs.device, dtype=torch.float32) for t in range(T): s = torch.einsum("bnij,bnj->bni", R[:, t], s) # per-block unit-norm projection: bounds growth (boosts amplify ‖s‖), # state lives on (S⁵)^nb. Matches the lattice's normalized recurrence. s = s / s.norm(dim=-1, keepdim=True).clamp_min(1e-6) S[:, t] = s return S def forward(self, idx, targets=None, ablate_scan=False): """ablate_scan=True → R→IDENTITY (state frozen at s0) but coefs/tower stay LIVE. This is the BYPASS TEST: if the model still predicts well with the rotors off, the grade tower is feeding operators to the readout past the state recurrence (a bypass) — the causal claim is dead. Correct pass: ablated PPL ≈ unigram floor. (Reported as absolute PPL, not a ratio — the ratio inflates mechanically as trained PPL drops.)""" B, T = idx.shape coefs = self.emit(idx) # (B,T,nb,15) S = self.scan(coefs, ablate=ablate_scan) # (B,T,nb,6) grade-1 if self.tower is not None: # causal shifts of the operator sequence (grade-2) B_prev = torch.zeros_like(coefs) B_prev[:, 1:] = coefs[:, :-1] B_prev2 = torch.zeros_like(coefs) B_prev2[:, 2:] = coefs[:, :-2] g3, g4, g5, g6 = self.tower(S, coefs, B_prev, B_prev2) feat = torch.cat([S, coefs, g3, g4, g5, g6], dim=-1) # (B,T,nb,63) feat = feat.reshape(B, T, -1) else: feat = S.reshape(B, T, -1) feat = self.state_norm(feat) logits = self.readout(feat) # (B,T,vocab) loss = None if targets is not None: loss = F.cross_entropy(logits.reshape(-1, logits.shape[-1]), targets.reshape(-1)) return logits, loss def param_count(self): return sum(p.numel() for p in self.parameters() if p.requires_grad) __all__ = ["OpEmitConfig", "OpEmitLM"]