Download t3v3_wedge_memory.py from mirrorethic/cl33-oplm: direct link, hf CLI and curl.
- Browser
- Download file 9.96 kB
-
https://huggingface.co/mirrorethic/cl33-oplm/resolve/3ee235d5c9521eb6993535736afc2d7e7758e3ca/t3v3_wedge_memory.py
- Command line
-
hf download hf://mirrorethic/cl33-oplm@3ee235d5c9521eb6993535736afc2d7e7758e3ca/t3v3_wedge_memory.py
-
curl -L -o t3v3_wedge_memory.py https://huggingface.co/mirrorethic/cl33-oplm/resolve/3ee235d5c9521eb6993535736afc2d7e7758e3ca/t3v3_wedge_memory.py
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) | |