Download model_v2.py from mirrorethic/cl33-oplm: direct link, hf CLI and curl.
- Browser
- Download file 20.9 kB
-
https://huggingface.co/mirrorethic/cl33-oplm/resolve/main/model_v2.py
- Command line
-
hf download hf://mirrorethic/cl33-oplm/model_v2.py
-
curl -L -o model_v2.py https://huggingface.co/mirrorethic/cl33-oplm/resolve/main/model_v2.py
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) | |
| 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)) | |
| 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"] | |