File size: 10,916 Bytes
3ee235d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 | """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"]
|