"""Reversible tape memory — forward-readable exact memory native to the algebra. See REVERSIBLE_TAPE_DESIGN.md (2026-07-09). Design principle: TOKEN CONTENT MAY SELECT; ONLY ALGEBRA MAY FLOW. - Addressing: context-free token keys (the wedge-validated mechanism) produce a causal softmax α over past tape positions. Token content touches ONLY α. - Read: r_t = Σ_i α_{t,i} · f(tape_i), where f yields pure algebra objects. The tape is the model's own reversible trace. The exact prefix products P_t = R_t···R_1 give the relative transport between any two positions: R_{t←i} = P_t · P_i⁻¹, with P_i⁻¹ = η P_iᵀ η (exact, no replay) Degeneracy (design §3): transported absolute states are trivial (R_{t←i}s_i ∝ s_t), so the non-degenerate per-position content is exactly: V1 "displacement": R_{t←i} applied to a fixed reference → trajectory geometry only V2 "state": s_i raw (frame-mismatched on purpose) → algebra-pure past state V3 "increment": B_i, the emitted bivector at i → the local operator event (token-decodability ~82% per reversibility.py M2 → carries the content MQAR needs, without token embeddings in the value path) Op-ablation behavior (transparency mechanics): zeroed operators ⇒ R=I ⇒ P_t=I, displacement=const, increment=0, state=s0 — every value mode collapses while the token keys survive. The read CONTENT is operator-derived by construction. """ from __future__ import annotations import torch import torch.nn as nn import torch.nn.functional as F VALUE_DIMS = {"displacement": 6, "state": 6, "increment": 15, "multi": 21} # displacement(6) + increment(15), one shared address class TapeMemory(nn.Module): """Content-addressed read over the reversible tape. O(T²) score like the main attention (fine at MQAR/32M scale; anchors+subsampling are the long-context lever, not needed here).""" def __init__(self, d_model: int, nb: int, value_mode: str = "increment", induction: bool = True, dual_address: bool = False, addr: str = "token"): super().__init__() assert value_mode in VALUE_DIMS, value_mode assert addr in ("token", "operator", "target"), addr self.nb, self.value_mode, self.induction, self.addr = nb, value_mode, induction, addr # context-free token keys — addressing power proven by the wedge token mode self.proj_q = nn.Linear(d_model, nb * 6) self.proj_k = nn.Linear(d_model, nb * 6) # v2.8 OPERATOR addressing (Garret): the address query is the LM's emitted OPERATOR, not the # token. Keys = past operators. "what stored operator-combination fits what I'm computing now" # instead of "what token matches" — the native relational address (op-dep memory), NOT the # redundant token channel that CE routes around (v2.6 co-option). Ablating operators kills BOTH # address and value → the whole tape becomes operator-derived (zero token bypass in addressing). if addr == "operator": self.proj_q_op = nn.Linear(nb * 15, nb * 6) self.proj_k_op = nn.Linear(nb * 15, nb * 6) # v2.8 TARGET-DRIVEN navigation (the diamond, facet 2): select by GOAL-vs-EFFECT fit, not # key-match. target = learned "what I want" from the current operator; offer = each record's # rotor ACTION on a learned reference (its composition effect). score = . if addr == "target": self.proj_target = nn.Linear(nb * 15, nb * 6) # goal, from the current operator Bq self.nav_ref = nn.Parameter(torch.randn(nb, 6) * 0.5) # reference the record's rotor acts on # v2.5 DUAL-ADDRESS (TAPE_ADDRESSING_V25.md): a token-FAITHFUL exact-match channel # alongside the learned proj address. Probes showed the learned proj gets co-opted # for LM context on natural text (attribution 0.98 raw-token → 0.006 learned); a # dedicated raw-token-cosine address recovers ~0.98. exact_gate zero-init → a fresh # v2.5 model is byte-identical to v2.4 at step 0, and v2.4 ckpts load unaffected. self.dual_address = dual_address if dual_address: self.exact_tau = nn.Parameter(torch.tensor(8.0)) # sharp init (blockade-like) self.exact_gate = nn.Parameter(torch.zeros(nb, 1)) # per-block, zero-init # V1's fixed reference direction per block (constant → carries no content; # the read then reflects trajectory displacement only). Only displacement- # bearing modes use it — don't create an unused Parameter otherwise. self.ref = (nn.Parameter(torch.randn(nb, 6) * 0.5) if value_mode in ("displacement", "multi") else None) def out_dim(self) -> int: return VALUE_DIMS[self.value_mode] def forward(self, S, Bs, R_state, tok_emb, eta, Bq=None, S_tape=None, Bs_tape=None, R_tape=None, tok_tape=None): """S (B,T,nb,6) states; Bs (B,T,nb,15) emitted bivectors; R_state (B,T,nb,6,6) per-step rotors; tok_emb (B,T,d_model) RAW token embedding (addressing only). Returns r (B,T,nb,out_dim), fp32. *_tape kwargs (SELF_STRUCTURE_CLAIMS.md, Claim-2 surgery): optional COUNTERFACTUAL tape record — the read head sees these instead of the factual history, while the live computation (scan/attention, and the queries q) stays factual. Keys k and values are drawn from the tape record (they are part of the record being counterfactually presented). Absent kwargs = factual behavior, byte-identical to before.""" te = tok_emb.float() B, T, nb = S.shape[:3] S = S.float() eta = eta.float() # counterfactual record substitution (read-side only) S_rec = S if S_tape is None else S_tape.float() Bs_rec = Bs if Bs_tape is None else Bs_tape R_rec = R_state if R_tape is None else R_tape te_rec = te if tok_tape is None else tok_tape.float() # --- addressing --- # queries from the FACTUAL present; keys from the (possibly counterfactual) tape record. # OPERATOR mode: address by the emitted operator (native relational key); token content # never touches the address. TOKEN mode: the original context-free token key. if self.addr == "target": # goal-vs-effect fit (navigation, NOT key-match; no induction shift — this is not adjacency). assert Bq is not None and R_rec is not None, "target addressing needs Bq and R_state (rotor of Bs)" target = self.proj_target(Bq.float().reshape(B, T, nb * 15)).view(B, T, nb, 6) # what I want ref = self.nav_ref / self.nav_ref.norm(dim=-1, keepdim=True).clamp_min(1e-6) # (nb,6) offer = torch.einsum("btnij,nj->btni", R_rec.float(), ref) # record EFFECT = rotor·ref (B,T,nb,6) scores = torch.einsum("btni,bsni->bnts", target, offer) / (6 ** 0.5) # else: if self.addr == "operator": assert Bq is not None, "operator addressing needs the query operator Bq" q = self.proj_q_op(Bq.float().reshape(B, T, nb * 15)).view(B, T, nb, 6) k = self.proj_k_op(Bs_rec.float().reshape(B, T, nb * 15)).view(B, T, nb, 6) else: q = self.proj_q(te).view(B, T, nb, 6) k = self.proj_k(te_rec).view(B, T, nb, 6) if self.induction: # key at i ← the PREVIOUS token: querying with a key-token selects the position of what # FOLLOWED it (MQAR adjacency, as in the wedge). Key-match only. k = torch.cat([torch.zeros_like(k[:, :1]), k[:, :-1]], dim=1) scores = torch.einsum("btni,bsni->bnts", q, k) / (6 ** 0.5) # (B,nb,T,T) causal = torch.triu(torch.ones(T, T, device=S.device, dtype=torch.bool), 0) scores = scores.masked_fill(causal, float("-inf")) # STRICT past (ibts", tq, tk) * self.exact_tau esc = esc.masked_fill(causal[None], float("-inf")) A_exact = torch.nan_to_num(F.softmax(esc, dim=-1), nan=0.0)[:, None] # (B,1,T,T) def _dual(r_ctx, value): if A_exact is None: return r_ctx r_ex = torch.einsum("bnts,bsni->btni", A_exact.expand(B, nb, T, T), value) return r_ctx + self.exact_gate * r_ex # gate zero-init → ≡ v2.4 # --- read (only algebra flows) --- if self.value_mode == "state": return _dual(torch.einsum("bnts,bsni->btni", A, S_rec), S_rec) if self.value_mode == "increment": return _dual(torch.einsum("bnts,bsni->btni", A, Bs_rec.float()), Bs_rec.float()) if self.value_mode == "multi": # one α, two reads: increment (15) + displacement (6) — the local-order # and global-path channels together (v2.4 design) r_inc = _dual(torch.einsum("bnts,bsni->btni", A, Bs_rec.float()), Bs_rec.float()) r_disp = self._displacement_read(A, R_rec.float(), S.device, B, T, nb, eta) return torch.cat([r_inc, r_disp], dim=-1) # displacement: r_t = P_t · Σ_i α_{t,i} (P_i⁻¹ · ref̂) return self._displacement_read(A, R_rec.float(), S.device, B, T, nb, eta) def _displacement_read(self, A, R, device, B, T, nb, eta): P = torch.empty(B, T, nb, 6, 6, device=device, dtype=torch.float32) acc = torch.eye(6, device=device).expand(B, nb, 6, 6).contiguous() for t in range(T): acc = R[:, t] @ acc P[:, t] = acc ref = self.ref / self.ref.norm(dim=-1, keepdim=True).clamp_min(1e-6) # (nb,6) # P_i⁻¹ = η P_iᵀ η Pinv = eta[None, None, None, :, None] * P.transpose(-1, -2) * eta[None, None, None, None, :] u = torch.einsum("btnij,nj->btni", Pinv, ref) # (B,T,nb,6) mix = torch.einsum("bnts,bsni->btni", A, u) return torch.einsum("btnij,btnj->btni", P, mix) __all__ = ["TapeMemory", "VALUE_DIMS"]