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