cl33-oplm / tape_memory.py
gman1911's picture
Selective reproducibility release for preprint v1.1 (One Object): two frozen 236M checkpoints, load-only model code, verified repro scripts, hashes
3ee235d verified
Raw History Blame
10.9 kB
"""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 = <target, offer>.
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) # <goal_t, effect_s>
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 (i<t)
A = F.softmax(scores, dim=-1)
A = torch.nan_to_num(A, nan=0.0) # t=0 has no past
# v2.5 DUAL-ADDRESS: token-FAITHFUL exact-match channel (raw token cosine, sharp).
# It only SELECTS (like A) — the value path is unchanged, so transparency holds.
A_exact = None
if self.dual_address:
tq = F.normalize(te, dim=-1) # (B,T,d) raw query token
tk = F.normalize(te_rec, dim=-1) # keys from the record
if self.induction:
tk = torch.cat([torch.zeros_like(tk[:, :1]), tk[:, :-1]], dim=1)
esc = torch.einsum("btd,bsd->bts", 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"]