cl33-oplm / model.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
7.39 kB
"""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"]