cl33-oplm / model_v2.py
gman1911's picture
Docstring precision: pinning cost = intervention artifact; LayerNorm = candidate mechanism (not isolated)
836287c verified
Raw History Blame Contribute Delete
20.9 kB
"""cl33-opLM v2 — operator-native attention over the state trajectory.
The v0/widen finding: the operator-only bottleneck forces a RECURRENT state that
compresses history; widening independent blocks doesn't beat compression. The
fix (Garret): restore attention, but map it 1:1 onto the algebra so it stays
operator-derived (causal transparency preserved).
Operator-native attention (per block = per head):
state scan: s_t = R(B_state_t) · s_{t-1} (as before, reversible)
query/key: Q_t = R(B_q_t) · s_t, K_s = R(B_k_s) · s_s (emitted rotors)
score: ⟨Q_t, K_s⟩_η (η-metric inner product, causal)
attend: a_t = Σ_s softmax(score)_ts · s_s (values = the states)
readout: LayerNorm(flatten a_t) → vocab
Because rotors preserve η, score(t,s) = ⟨s_t, (R_q⁻¹R_k)·s_s⟩_η — attention is
the alignment of the current state with a LEARNED-RELATIVELY-ROTATED past state.
Everything is operator-derived → operator-only bottleneck holds. STALE-PREDICTION
NOTE (2026-09-10, kept for the record): "the R_state→identity bypass collapses to
unigram" was written for the PRE-TOWER readout and was true of it; with the grade
tower (shipped config) the readout has a direct route to B_t..B_{t-2} and identity-
scan costs only ~1.2x (an intervention artifact, not lost history — time-shuffling S
costs ~1.0x; shared-LayerNorm shift is the candidate mechanism, not yet isolated;
measured internally and independently replicated by N. Watson 2026-09-10).
At LM inference the transported state is causally decorative; in sole-channel/tape
configs it is load-bearing. See paper §3.2.
Scan + attention in fp32 (bf16 proven-unsafe on the recurrence).
"""
from __future__ import annotations
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, ETA # type: ignore
from model import EmitterBlock # reuse the emitter block
from wedge import GradeTower # grade tower for short-context (current operator)
from t3v3_wedge_memory import WedgeMemory, TokenCopyMemory # associative / copy memory
from tape_memory import TapeMemory # reversible-tape read (address by token, flow algebra)
@dataclass
class OpEmitV2Config:
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 = 32 # = attention heads (each attends in its SO(3,3))
coef_clip: float = 1.0
op_init_scale: float = 0.02
qk_init_scale: float = 0.1 # q/k rotors can be larger (learn what to attend)
use_grade_tower: bool = True # v2.2: current-operator grade tower in readout
use_wedge_memory: bool = False # wedge bivector associative memory (KV binding)
wedge_key_source: str = "operator" # "operator" | "state" | "token"
wedge_value_source: str = "same" # "same" | "operator" (G1: algebra value w/ token key)
wedge_delta_rule: bool = False # DeltaNet residual write (D3 interference fix)
use_token_copy: bool = False # token-content copy channel (measured transparency cost)
use_tape_memory: bool = False # reversible-tape read (REVERSIBLE_TAPE_DESIGN.md)
tape_value_mode: str = "increment" # "displacement" | "state" | "increment"
tape_dual_address: bool = False # v2.5: + token-faithful exact-match channel (TAPE_ADDRESSING_V25.md)
tape_addr: str = "token" # v2.8: "token" (address by token embedding) | "operator" (address
# by the emitted LM operator — native relational key, not the
# redundant token channel CE routes around). Pairs with tape_compose.
tape_compose: bool = False # v2.8: gp-COMPOSE recalled operator with the query rotor before
# readout (genesis marriage compose + v2.7 product-with-query). The
# recalled op stops being read in isolation — it composes with the
# current computation. Needs tape_value_mode=increment (15-d recall).
tape_compose_mode: str = "add" # "add" = additive zero-init logit (SAFE, out of LN; augments the
# distribution). "tilt" = recalled operator (gated, zero-init) TILTS
# the decoded state itself, decoded by the normal readout (augments
# BEHAVIOR; in the LN → tests whether it survives tape-silencing).
scan_only: bool = False # zero the O(T²) attention — tape/scan carry history
dropout: float = 0.0
grad_checkpoint: bool = False # activation-checkpoint the emitter transformer blocks: recompute
# in backward instead of storing (bit-identical forward, unchanged
# architecture/reversibility). Frees the dominant activation memory
# → larger batch → better sequential-scan SM occupancy on small GPUs.
class OpEmitLMv2(nn.Module):
def __init__(self, c: OpEmitV2Config):
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)
# emit 3 bivectors per block: state-evolution, query-rotor, key-rotor
self.op_head = nn.Linear(c.d_model, c.n_blocks * 3 * N_GEN)
nn.init.normal_(self.op_head.weight, std=1e-3)
nn.init.zeros_(self.op_head.bias)
self.s0 = nn.Parameter(torch.randn(c.n_blocks, 6) * 0.5)
# readout features per block:
# attended(6) + s_t(6=g1) [residual: long-context attn + current state]
# + grade tower g2(15)+g3(20)+g4(15)+g5(6)+g6(1)=57 [v2.2: current
# operator, rich at pos 1 — fixes the rotation-of-s0 short-context
# limit]. All operator-derived → bottleneck preserved.
self.tower = GradeTower() if c.use_grade_tower else None
self.wedge = (WedgeMemory(key_source=getattr(c, "wedge_key_source", "operator"),
value_source=getattr(c, "wedge_value_source", "same"),
delta_rule=getattr(c, "wedge_delta_rule", False),
d_model=c.d_model, nb=c.n_blocks)
if getattr(c, "use_wedge_memory", False) else None)
self.token_copy = (TokenCopyMemory(c.d_model, c.n_blocks)
if getattr(c, "use_token_copy", False) else None)
self.tape = (TapeMemory(c.d_model, c.n_blocks,
value_mode=getattr(c, "tape_value_mode", "increment"),
dual_address=getattr(c, "tape_dual_address", False),
addr=getattr(c, "tape_addr", "token"))
if getattr(c, "use_tape_memory", False) else None)
# readout features per block: +6 each for the wedge read and the copy read.
# The TAPE is deliberately NOT in this vector — it is a SEPARATE post-LayerNorm additive
# readout term (see assemble). Reason: concatenating the tape into the shared LayerNorm lets
# any tape contribution perturb the normalization of the base features, so the optimizer
# silences the tape to protect the base LM (observed: v2.4 gate 0.01→0.0005). Keeping the
# tape OUT of the LN leaves the base bit-exact (clean warm-start) and lets the tape co-train.
feat_per_block = (12 + (6 if self.wedge is not None else 0)
+ (6 if self.token_copy is not None else 0)
+ (57 if c.use_grade_tower else 0))
self.state_norm = nn.LayerNorm(c.n_blocks * feat_per_block)
self.readout = nn.Linear(c.n_blocks * feat_per_block, c.vocab_size, bias=False)
# tape = additive side channel, ZERO-INIT readout → off at step 0, learns on. The zero-init
# readout IS the clean off-switch (no gate scalar needed, no LN-suppression dynamic).
if self.tape is not None:
self.tape_readout = nn.Linear(self.tape.out_dim() * c.n_blocks, c.vocab_size, bias=False)
nn.init.zeros_(self.tape_readout.weight)
# v2.8 compose channel: recalled operator gp-composed with the query rotor, applied to state.
# Separate ZERO-INIT readout (same off-switch pattern) → step-0 ≡ base, co-trains on.
self.tape_compose = getattr(c, "tape_compose", False) and self.tape is not None
if self.tape_compose:
assert self.tape.out_dim() == 15, "tape_compose needs 15-d bivector recall (tape_value_mode=increment)"
if getattr(c, "tape_compose_mode", "add") == "tilt":
# gated recalled operator tilts the decoded state; per-block gate ZERO-INIT → R_tilt=I,
# step-0 state unchanged. In the LN via parts=[a, S_tilt] → the behavior-augmenting arm.
self.tape_tilt_gate = nn.Parameter(torch.zeros(c.n_blocks, 1))
else:
self.tape_compose_readout = nn.Linear(c.n_blocks * 6, c.vocab_size, bias=False)
nn.init.zeros_(self.tape_compose_readout.weight)
self.register_buffer("G", build_generators(dtype=torch.float32), persistent=False)
self.register_buffer("eta", ETA.clone(), 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]
if getattr(self.c, "grad_checkpoint", False) and self.training:
from torch.utils.checkpoint import checkpoint
for blk in self.blocks:
x = checkpoint(blk, x, use_reentrant=False) # recompute in backward, free activations
else:
for blk in self.blocks:
x = blk(x)
h = self.norm(x)
raw = self.op_head(h).view(B, T, self.c.n_blocks, 3, N_GEN)
Bs, Bq, Bk = raw.unbind(-2) # each (B,T,nb,15)
def clip(bv, scale):
bv = bv * scale
n = bv.norm(dim=-1, keepdim=True)
return bv * (self.c.coef_clip / n.clamp_min(self.c.coef_clip))
return (clip(Bs, self.c.op_init_scale),
clip(Bq, self.c.qk_init_scale),
clip(Bk, self.c.qk_init_scale))
@torch.no_grad()
def emitter_hidden(self, idx):
"""The emitter's final hidden h_t (baseline for recoverability probes)."""
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)
return self.norm(x) # (B,T,d_model)
def scan(self, Bs, ablate=False):
"""Reversible state recurrence. Returns S (B,T,nb,6), fp32."""
B, T, nb, _ = Bs.shape
Bs = Bs.float()
R = (torch.eye(6, device=Bs.device).expand(B, T, nb, 6, 6) if ablate
else rotor_from_coefs(Bs, self.G.float()))
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=Bs.device, dtype=torch.float32)
for t in range(T):
s = torch.einsum("bnij,bnj->bni", R[:, t], s)
s = s / s.norm(dim=-1, keepdim=True).clamp_min(1e-6)
S[:, t] = s
return S
def scan_parallel(self, Bs, ablate=False):
"""Parallel associative scan (PARALLEL_SCAN_SPEC): the per-step normalize cancels, so
S_t = normalize(P_t·s0) with P_t = R_t···R_0 a prefix product. Hillis-Steele scan under the
associative operator A∘B = normalize_F(A·B) (Frobenius-normalized to avoid boost overflow).
O(log T) depth. GATED against scan() — must match bit-for-bit before it replaces it."""
B, T, nb, _ = Bs.shape
Bs = Bs.float()
R = (torch.eye(6, device=Bs.device).expand(B, T, nb, 6, 6).contiguous() if ablate
else rotor_from_coefs(Bs, self.G.float()))
def nF(M):
return M / (M.reshape(*M.shape[:-2], 36).norm(dim=-1)[..., None, None] + 1e-30)
P = nF(R) # (B,T,nb,6,6)
I = torch.eye(6, device=Bs.device, dtype=P.dtype).expand(B, 1, nb, 6, 6)
idx = torch.arange(T, device=Bs.device)
d = 1
while d < T:
right = torch.cat([I.expand(B, d, nb, 6, 6), P[:, :T - d]], dim=1) # right[t]=P[t-d], I for t<d
combined = nF(P @ right) # newest(left) @ older(right), correct order
mask = (idx >= d)[None, :, None, None, None]
P = torch.where(mask, combined, P)
d *= 2
s0 = self.s0.float().expand(B, nb, 6)
s0 = s0 / s0.norm(dim=-1, keepdim=True).clamp_min(1e-6)
S = torch.einsum("btnij,bnj->btni", P, s0)
return S / S.norm(dim=-1, keepdim=True).clamp_min(1e-6)
def attend(self, S, Bq, Bk):
"""Operator-native causal attention over the state trajectory.
S (B,T,nb,6); Bq,Bk (B,T,nb,15). Returns attended (B,T,nb,6)."""
B, T, nb, _ = S.shape
Rq = rotor_from_coefs(Bq.float(), self.G.float()) # (B,T,nb,6,6)
Rk = rotor_from_coefs(Bk.float(), self.G.float())
Q = torch.einsum("btnij,btnj->btni", Rq, S) # (B,T,nb,6)
K = torch.einsum("btnij,btnj->btni", Rk, S)
# η-metric scores: ⟨Q_t, K_s⟩_η, per block/head. (B,nb,T,T)
Kw = K * self.eta.to(K.dtype) # apply η to keys
scores = torch.einsum("btni,bsni->bnts", Q, Kw) / (6 ** 0.5)
causal = torch.triu(torch.ones(T, T, device=S.device, dtype=torch.bool), 1)
scores = scores.masked_fill(causal, float("-inf"))
A = F.softmax(scores, dim=-1) # (B,nb,T,T)
attended = torch.einsum("bnts,bsni->btni", A, S) # values = states
return attended
def assemble(self, Bs, Bq, Bk, scan_only=False, tok_emb=None):
"""Full pass from emitted operators → logits (scan + attention + tower
+ readout). Separated from emit() so CONTROL interventions can perturb
the operators and re-run only the downstream. Returns (B,T,vocab).
scan_only=True zeroes the cross-position attention path, so recall must
come from the recurrent reversible state alone (the fair vs-xLSTM memory
test — isolates the linear-recurrent memory from the O(T²) attention)."""
B, T = Bs.shape[:2]
S = self.scan(Bs, ablate=False) # (B,T,nb,6)
# tape read (computed here so a 'tilt' compose can act on the decoded state below)
tape_read = None
if self.tape is not None and tok_emb is not None:
R_state = (rotor_from_coefs(Bs.float(), self.G.float())
if (self.tape.value_mode in ("displacement", "multi")
or getattr(self.tape, "addr", "token") == "target") else None)
tape_read = self.tape(S, Bs, R_state, tok_emb, self.eta, Bq=Bq) # (B,T,nb,out_dim)
# v2.8 tilt-compose: the recalled operator (gated, zero-init) tilts the state that gets decoded
# by the NORMAL readout — augments behavior, then decoded like usual.
S_ro = S
if (tape_read is not None and getattr(self, "tape_compose", False)
and getattr(self.c, "tape_compose_mode", "add") == "tilt"):
tilt_biv = self.tape_tilt_gate * tape_read.float() # (B,T,nb,15); gate 0 → identity
R_tilt = rotor_from_coefs(tilt_biv, self.G.float()) # (B,T,nb,6,6)
S_ro = torch.einsum("btnij,btnj->btni", R_tilt, S) # recalled op tilts the decoded state
a = torch.zeros_like(S) if scan_only else self.attend(S, Bq, Bk)
parts = [a, S_ro]
if self.wedge is not None:
# O(T) linear-recurrent associative memory. operator mode keys on the
# per-token emitted operators (content-bearing); state mode (ablation)
# keys on the transported state + adjoint-transports the memory.
Rq = Rk = R_state = None
if self.wedge.key_source == "state":
G32 = self.G.float()
Rq = rotor_from_coefs(Bq.float(), G32)
Rk = rotor_from_coefs(Bk.float(), G32)
R_state = rotor_from_coefs(Bs.float(), G32)
parts.append(self.wedge(S, Rq, Rk, self.eta, R_state=R_state,
Bs=Bs, Bq=Bq, Bk=Bk, tok_emb=tok_emb)) # r_t (B,T,nb,6)
if self.token_copy is not None and tok_emb is not None:
# token-content copy channel (bypasses operators; transparency cost measured)
parts.append(self.token_copy(tok_emb)) # (B,T,nb,6)
# (tape_read computed above, before the state-tilt; consumed as a SEPARATE post-LN additive
# term below in "add" mode, or already applied as a state tilt above in "tilt" mode.)
if self.tower is not None:
Bs_prev = torch.zeros_like(Bs); Bs_prev[:, 1:] = Bs[:, :-1]
Bs_prev2 = torch.zeros_like(Bs); Bs_prev2[:, 2:] = Bs[:, :-2]
g3, g4, g5, g6 = self.tower(S, Bs, Bs_prev, Bs_prev2)
parts += [Bs, g3, g4, g5, g6]
feat = self.state_norm(torch.cat(parts, dim=-1).reshape(B, T, -1))
logits = self.readout(feat)
if tape_read is not None: # additive, post-LN, zero-init
logits = logits + self.tape_readout(tape_read.reshape(B, T, -1).to(logits.dtype))
if (tape_read is not None and getattr(self, "tape_compose", False)
and getattr(self.c, "tape_compose_mode", "add") == "add"):
# gp-COMPOSE the recalled operator with the current query rotor (exact matrix/rotor
# composition — lossless in the grade-1 action), apply to the state. The recalled
# operation now composes with the current computation instead of being read in isolation:
# genesis-marriage compose + v2.7's product-with-query. Zero-init readout → off at step 0.
G32 = self.G.float()
R_rec = rotor_from_coefs(tape_read.float(), G32) # (B,T,nb,6,6) recalled operator
R_q = rotor_from_coefs(Bq.float(), G32) # (B,T,nb,6,6) current query
R_comp = torch.einsum("btnij,btnjk->btnik", R_rec, R_q) # exact rotor composition
comp = torch.einsum("btnij,btnj->btni", R_comp, S) # composed op applied to state
logits = logits + self.tape_compose_readout(comp.reshape(B, T, -1).to(logits.dtype))
return logits
def forward(self, idx, targets=None, ablate_scan=False, amp_emit=False):
# amp_emit: run the transformer emitter in fp16 (tensor-core speedup, safe)
# but the operator algebra (scan/attend/tower/readout) in fp32 — fp16 there
# overflows (η-metric scores + nested wedge products) → NaN. Split keeps the
# bulk of compute fast while the delicate algebra stays numerically exact.
if amp_emit:
with torch.autocast("cuda", dtype=torch.float16):
Bs, Bq, Bk = self.emit(idx)
Bs, Bq, Bk = Bs.float(), Bq.float(), Bk.float()
else:
Bs, Bq, Bk = self.emit(idx)
if ablate_scan:
Bs = torch.zeros_like(Bs) # transparency test
# tok_emb (raw embedding) feeds the token-copy / token-key channels; on the
# ablate path it is retained → bypass ratio then reflects how much those
# channels carry prediction around the (zeroed) operators.
tok_emb = self.tok_embed(idx) if (self.token_copy is not None
or self.tape is not None
or (self.wedge is not None and self.wedge.key_source == "token")) else None
logits = self.assemble(Bs, Bq, Bk, scan_only=getattr(self.c, "scan_only", False),
tok_emb=tok_emb)
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__ = ["OpEmitV2Config", "OpEmitLMv2"]