cl33-oplm / t3v3_wedge_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
9.96 kB
"""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)