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