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