"""Wedge associative memory — the geometric-algebra-native fix for cl33-opLM's key→value binding failure (MQAR recall ≈ 1/KV). See WEDGE_MEMORY_DESIGN.md. A rotor can only rotate the state; it cannot do an outer-product WRITE, which is how associative binding is stored (mLSTM: C += v kᵀ, read C q). The outer product in geometric algebra is the wedge, so we add a persistent leaky bivector memory: M_t = γ_t · M_{t-1} + (K_t ∧ V_t) (write; bivector, per block) r_t = Q_t ⌋ M_{t-1} (read, before write — causal, no self-match) With the Cl(3,3) η-metric, the wedge/contraction reduce to a clean matrix form: K∧V as matrix: W = (V Kᵀ − K Vᵀ) · diag(η) so that Q ⌋ (K∧V) = W q = ⟨Q,K⟩_η V − ⟨Q,V⟩_η K i.e. when the query matches a stored key, the read returns that key's value. This is fast-weight / linear-attention associative memory, native to the algebra — the O(T) linear-recurrent sibling of the O(T²) operator-attention (same emitted Q/K). """ from __future__ import annotations import torch import torch.nn as nn class WedgeMemory(nn.Module): """Per-block leaky bivector associative memory. Reuses the emitted query/key rotors (via Rq, Rk applied to the state) — no new emission. Returns the read r_t (grade-1, 6-d per block) to concatenate into the readout features. Value = the transported state s_t (algebra-pure; keeps the operator-only bottleneck / transparency — see design §5/§6 for the richer-V escalation).""" def __init__(self, gate_bias_init: float = 3.0, transport: bool = True, key_source: str = "operator", n_gen: int = 15, induction: bool = True, d_model: int = None, nb: int = None, value_source: str = "same", delta_rule: bool = False): super().__init__() self.induction = induction # key ← PREVIOUS token/operator (adjacency for MQAR) self.key_source = key_source # value_source: "same" = value from the same source as the key (default, # legacy). "operator" = value is a learned 6-d projection of the emitted # state-operator Bs (algebra-only, ablatable) — the G1 experiment: token KEY # (clean context-free address) + ALGEBRA VALUE (LINEAR_TRANSPARENT_MEMORY.md). self.value_source = value_source self.delta_rule = delta_rule # DeltaNet residual write: store v - M·k (reduces # write interference by construction) — G/D3. if value_source == "operator": self.proj_v_op = nn.Linear(n_gen, 6) self.nb = nb # input-dependent forget gate γ_t = σ(w·x + b). bias_init 3.0 → γ ≈ 0.95. # key_source options: # "operator": Q/K/V from emitted operators B_q/B_k/B_state (per-token but # CONTEXT-MIXED by the causal emitter → induction can't form a # clean key; falsified, recall ≈ 1/KV). # "state": from the transported state (frame-entangled; + adjoint transport). # "token": grade-1 Cl(3,3) projection of the RAW token embedding — # CONTEXT-FREE → clean induction key. Fully-transparent # (algebra-native η-wedge read) version of the token-copy win. if key_source == "operator": self.proj_q = nn.Linear(n_gen, 6); self.proj_k = nn.Linear(n_gen, 6) self.proj_v = nn.Linear(n_gen, 6) gate_in = n_gen elif key_source == "token": self.proj_q = nn.Linear(d_model, nb * 6); self.proj_k = nn.Linear(d_model, nb * 6) self.proj_v = nn.Linear(d_model, nb * 6) gate_in = d_model else: gate_in = 6 self.gate = nn.Linear(gate_in, 1) nn.init.zeros_(self.gate.weight) nn.init.constant_(self.gate.bias, gate_bias_init) self.transport = transport and key_source == "state" def forward(self, S, Rq, Rk, eta, R_state=None, Bs=None, Bq=None, Bk=None, tok_emb=None): """Returns r (B,T,nb,6). fp32. operator: Bq/Bk/Bs (B,T,nb,15); state: S/Rq/Rk; token: tok_emb (B,T,d_model) → grade-1 projections.""" eta = eta.float() B, T, nb = S.shape[:3] if self.key_source == "operator": Q = self.proj_q(Bq.float()); K = self.proj_k(Bk.float()) V = self.proj_v(Bs.float()) # per-token, context-mixed gamma = torch.sigmoid(self.gate(Bs.float())).squeeze(-1) elif self.key_source == "token": te = tok_emb.float(); nb = self.nb Q = self.proj_q(te).view(B, T, nb, 6) # context-free grade-1 keys K = self.proj_k(te).view(B, T, nb, 6) if self.value_source == "operator": # G1: algebra-only value — a learned projection of the emitted # state-operator. Ablating operators (Bs=0) zeroes the value content # → read dies → zero bypass (the property token-value lacked). V = self.proj_v_op(Bs.float()).view(B, T, nb, 6) else: V = self.proj_v(te).view(B, T, nb, 6) # token content (bypass) gamma = torch.sigmoid(self.gate(te)).expand(B, T, nb) else: S = S.float(); Rq = Rq.float(); Rk = Rk.float() Q = torch.einsum("btnij,btnj->btni", Rq, S) K = torch.einsum("btnij,btnj->btni", Rk, S) V = S gamma = torch.sigmoid(self.gate(S)).squeeze(-1) if self.induction: # key at t ← the PREVIOUS token/operator → value_t stored under its # predecessor, so a query retrieves what FOLLOWED it (MQAR adjacency). K = torch.cat([torch.zeros_like(K[:, :1]), K[:, :-1]], dim=1) etad = eta[None, None, :] # (1,1,6) for R⁻¹ = ηRᵀη if self.transport and R_state is not None: R_state = R_state.float() M = torch.zeros(B, nb, 6, 6, device=S.device, dtype=torch.float32) reads = [] for t in range(T): if self.transport and R_state is not None: Rt = R_state[:, t] Rt_inv = etad[..., None] * Rt.transpose(-1, -2) * etad[..., None, :] M = Rt @ M @ Rt_inv q_t = Q[:, t] # (B,nb,6) # READ before write (causal; avoids the query matching its own write) reads.append(torch.einsum("bnij,bnj->bni", M, q_t)) # WRITE: W = (v_eff K_tᵀ − K_t v_effᵀ)·diag(η) = the bivector k∧v_eff as a matrix v_t, k_t = V[:, t], K[:, t] if self.delta_rule: # DeltaNet residual: subtract what M already returns for THIS key, so the # write corrects rather than clobbers (interference reduction by construction). # STABILIZERS (2026-07-11): L2-normalize key (bounds ||M.k||) + beta write-gate # in (0,1) => contraction guaranteed, fixes the M-blowup NaN. k_t = k_t / (k_t.norm(dim=-1, keepdim=True) + 1e-6) vpred = torch.einsum("bnij,bnj->bni", M, k_t) v_eff = 0.5 * (v_t - vpred) else: v_eff = v_t W = (torch.einsum("bni,bnj->bnij", v_eff, k_t) - torch.einsum("bni,bnj->bnij", k_t, v_eff)) * eta[None, None, None, :] g = gamma[:, t][..., None, None] # (B,nb,1,1) M = g * M + W return torch.stack(reads, dim=1) # (B,T,nb,6) class TokenCopyMemory(nn.Module): """Minimal token-content value/copy channel — the measured transparency cost of KV recall. A per-block fast-weight memory keyed on the RAW token embedding (context-free AND position-free → the same token is the same key everywhere, so it is matchable across positions and the value TOKEN can be copied). This deliberately reads token content that bypasses the operator path; how much the model leans on it (vs the operators) is the transparency price, measured directly. Plain outer-product fast weights (token embeddings are not in the Cl(3,3) metric).""" def __init__(self, d_model: int, nb: int, key_dim: int = 6, gate_bias_init: float = 3.0, induction: bool = True): super().__init__() self.nb, self.kd = nb, key_dim self.induction = induction # key ← PREVIOUS token (the shift that makes MQAR # solvable: store value under the key that preceded it) self.qkv = nn.Linear(d_model, nb * key_dim * 3) self.gate = nn.Linear(d_model, nb) nn.init.constant_(self.gate.bias, gate_bias_init) def forward(self, tok_emb): """tok_emb (B,T,d_model) — RAW token embedding (no pos). Returns r (B,T,nb,kd).""" tok_emb = tok_emb.float() B, T, _ = tok_emb.shape qkv = self.qkv(tok_emb).view(B, T, self.nb, 3, self.kd) q, k, v = qkv[..., 0, :], qkv[..., 1, :], qkv[..., 2, :] # (B,T,nb,kd) if self.induction: # key at t ← the PREVIOUS token → value_t is stored under its predecessor, # so querying with a token retrieves what FOLLOWED it (induction head). k = torch.cat([torch.zeros_like(k[:, :1]), k[:, :-1]], dim=1) gamma = torch.sigmoid(self.gate(tok_emb)) # (B,T,nb) M = torch.zeros(B, self.nb, self.kd, self.kd, device=tok_emb.device, dtype=torch.float32) reads = [] for t in range(T): reads.append(torch.einsum("bnij,bnj->bni", M, q[:, t])) # read before write outer = torch.einsum("bni,bnj->bnij", v[:, t], k[:, t]) # v ⊗ k M = gamma[:, t][..., None, None] * M + outer return torch.stack(reads, dim=1) # (B,T,nb,kd)