File size: 7,389 Bytes
3ee235d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 | """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"]
|