Selective reproducibility release for preprint v1.1 (One Object): two frozen 236M checkpoints, load-only model code, verified repro scripts, hashes
Browse files- README.md +68 -0
- REPRODUCE.md +68 -0
- SHA256SUMS +11 -0
- cl33_oplm_chat_236m.pt +3 -0
- cl33_oplm_prose_236m.pt +3 -0
- invert_probe_chat.pt +3 -0
- model.py +165 -0
- model_v2.py +332 -0
- repro_bottleneck.py +59 -0
- repro_reverse_readout.py +47 -0
- so33.py +66 -0
- t3v3_wedge_memory.py +173 -0
- tape_memory.py +180 -0
- wedge.py +72 -0
README.md
ADDED
|
@@ -0,0 +1,68 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# cl33-opLM — selective reproducibility release (preprint v1.1)
|
| 2 |
+
|
| 3 |
+
Paper: **One Object: Memory, Navigation, and Reportability in an Operator-Only
|
| 4 |
+
Language Model** — https://t3atlas.dev/cl33/paper/ · Live demo: https://cl33.t3atlas.dev
|
| 5 |
+
Author: Garret Sutherland, MirrorEthic LLC.
|
| 6 |
+
|
| 7 |
+
This bundle contains what is necessary to **independently test the published claims**
|
| 8 |
+
on frozen artifacts. It is deliberately not the training stack: the paper's §12
|
| 9 |
+
program is ongoing and its machinery is not included. Reproducibility surface ≠
|
| 10 |
+
complete source disclosure.
|
| 11 |
+
|
| 12 |
+
## Contents
|
| 13 |
+
|
| 14 |
+
| file | what | sha256 |
|
| 15 |
+
|---|---|---|
|
| 16 |
+
| `cl33_oplm_prose_236m.pt` | prose base, 236.5M, step 189307 (Table 1b checkpoint) | `fe407328…6a1fbf` |
|
| 17 |
+
| `cl33_oplm_chat_236m.pt` | chat/serving model, step 13996 (the cl33.t3atlas.dev model) | `ddd042a6…559535` |
|
| 18 |
+
| `invert_probe_chat.pt` | reverse-readout probe (held-out top-1 0.860, card inside) | `91201342…ca79ef` |
|
| 19 |
+
| `model_v2.py` + `model.py` + `so33.py` + `wedge.py` + `t3v3_wedge_memory.py` + `tape_memory.py` | model definition (load-only) | — |
|
| 20 |
+
| `repro_bottleneck.py` | Claim 1: the mandatory operator bottleneck | — |
|
| 21 |
+
| `repro_reverse_readout.py`| Claim 2: the operator stream is a transcript | — |
|
| 22 |
+
| `SHA256SUMS` | full hashes | — |
|
| 23 |
+
|
| 24 |
+
Full hashes in `SHA256SUMS`. Deps: `torch`, `transformers`, `datasets` (Python ≥3.10).
|
| 25 |
+
|
| 26 |
+
## Claim 1 — the operator bottleneck is mandatory
|
| 27 |
+
|
| 28 |
+
Zero the emitted operators; the model loses its only path to output.
|
| 29 |
+
|
| 30 |
+
```
|
| 31 |
+
python repro_bottleneck.py --ckpt cl33_oplm_prose_236m.pt
|
| 32 |
+
python repro_bottleneck.py --ckpt cl33_oplm_chat_236m.pt
|
| 33 |
+
```
|
| 34 |
+
|
| 35 |
+
Expected (WikiText-103 test, public data, seq 1024 — measured on this exact bundle):
|
| 36 |
+
|
| 37 |
+
| ckpt | native PPL | ops-off PPL | ratio |
|
| 38 |
+
|---|---|---|---|
|
| 39 |
+
| prose | ≈61 | ≈6,500 | **≈106×** |
|
| 40 |
+
| chat | ≈186 | ≈21,000 | **≈112×** |
|
| 41 |
+
|
| 42 |
+
(The paper's 314× is the chat checkpoint on its in-domain validation mix; the claim
|
| 43 |
+
is the order of magnitude, and it holds off-domain.)
|
| 44 |
+
|
| 45 |
+
## Claim 2 — the operator stream is a readable transcript
|
| 46 |
+
|
| 47 |
+
A probe that sees ONLY the emitted operators (no token input) decodes the text:
|
| 48 |
+
|
| 49 |
+
```
|
| 50 |
+
python repro_reverse_readout.py --text "any sentence you like"
|
| 51 |
+
```
|
| 52 |
+
|
| 53 |
+
Expected: ~0.86 top-1 on typical English (the probe's held-out card prints on load;
|
| 54 |
+
rare words fail toward semantic neighbors — that is the paper's §7 claim, not a bug).
|
| 55 |
+
|
| 56 |
+
## What is NOT here, and why
|
| 57 |
+
|
| 58 |
+
Training orchestration, data pipelines, the §12 memory-organ program (labeled ongoing
|
| 59 |
+
in the paper), and downstream control/steering machinery. The claims those support are
|
| 60 |
+
either reported with their own dated work-log provenance (paper, Appendix R) or not
|
| 61 |
+
yet published. This bundle is scoped to verify what the preprint asserts about these
|
| 62 |
+
frozen artifacts.
|
| 63 |
+
|
| 64 |
+
## Checksums / provenance
|
| 65 |
+
|
| 66 |
+
Both checkpoints are weights-only exports (optimizer state stripped) of the exact
|
| 67 |
+
training checkpoints named in the paper's Appendix R. Verify with:
|
| 68 |
+
`sha256sum -c SHA256SUMS`
|
REPRODUCE.md
ADDED
|
@@ -0,0 +1,68 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# cl33-opLM — selective reproducibility release (preprint v1.1)
|
| 2 |
+
|
| 3 |
+
Paper: **One Object: Memory, Navigation, and Reportability in an Operator-Only
|
| 4 |
+
Language Model** — https://t3atlas.dev/cl33/paper/ · Live demo: https://cl33.t3atlas.dev
|
| 5 |
+
Author: Garret Sutherland, MirrorEthic LLC.
|
| 6 |
+
|
| 7 |
+
This bundle contains what is necessary to **independently test the published claims**
|
| 8 |
+
on frozen artifacts. It is deliberately not the training stack: the paper's §12
|
| 9 |
+
program is ongoing and its machinery is not included. Reproducibility surface ≠
|
| 10 |
+
complete source disclosure.
|
| 11 |
+
|
| 12 |
+
## Contents
|
| 13 |
+
|
| 14 |
+
| file | what | sha256 |
|
| 15 |
+
|---|---|---|
|
| 16 |
+
| `cl33_oplm_prose_236m.pt` | prose base, 236.5M, step 189307 (Table 1b checkpoint) | `fe407328…6a1fbf` |
|
| 17 |
+
| `cl33_oplm_chat_236m.pt` | chat/serving model, step 13996 (the cl33.t3atlas.dev model) | `ddd042a6…559535` |
|
| 18 |
+
| `invert_probe_chat.pt` | reverse-readout probe (held-out top-1 0.860, card inside) | `91201342…ca79ef` |
|
| 19 |
+
| `model_v2.py` + `model.py` + `so33.py` + `wedge.py` + `t3v3_wedge_memory.py` + `tape_memory.py` | model definition (load-only) | — |
|
| 20 |
+
| `repro_bottleneck.py` | Claim 1: the mandatory operator bottleneck | — |
|
| 21 |
+
| `repro_reverse_readout.py`| Claim 2: the operator stream is a transcript | — |
|
| 22 |
+
| `SHA256SUMS` | full hashes | — |
|
| 23 |
+
|
| 24 |
+
Full hashes in `SHA256SUMS`. Deps: `torch`, `transformers`, `datasets` (Python ≥3.10).
|
| 25 |
+
|
| 26 |
+
## Claim 1 — the operator bottleneck is mandatory
|
| 27 |
+
|
| 28 |
+
Zero the emitted operators; the model loses its only path to output.
|
| 29 |
+
|
| 30 |
+
```
|
| 31 |
+
python repro_bottleneck.py --ckpt cl33_oplm_prose_236m.pt
|
| 32 |
+
python repro_bottleneck.py --ckpt cl33_oplm_chat_236m.pt
|
| 33 |
+
```
|
| 34 |
+
|
| 35 |
+
Expected (WikiText-103 test, public data, seq 1024 — measured on this exact bundle):
|
| 36 |
+
|
| 37 |
+
| ckpt | native PPL | ops-off PPL | ratio |
|
| 38 |
+
|---|---|---|---|
|
| 39 |
+
| prose | ≈61 | ≈6,500 | **≈106×** |
|
| 40 |
+
| chat | ≈186 | ≈21,000 | **≈112×** |
|
| 41 |
+
|
| 42 |
+
(The paper's 314× is the chat checkpoint on its in-domain validation mix; the claim
|
| 43 |
+
is the order of magnitude, and it holds off-domain.)
|
| 44 |
+
|
| 45 |
+
## Claim 2 — the operator stream is a readable transcript
|
| 46 |
+
|
| 47 |
+
A probe that sees ONLY the emitted operators (no token input) decodes the text:
|
| 48 |
+
|
| 49 |
+
```
|
| 50 |
+
python repro_reverse_readout.py --text "any sentence you like"
|
| 51 |
+
```
|
| 52 |
+
|
| 53 |
+
Expected: ~0.86 top-1 on typical English (the probe's held-out card prints on load;
|
| 54 |
+
rare words fail toward semantic neighbors — that is the paper's §7 claim, not a bug).
|
| 55 |
+
|
| 56 |
+
## What is NOT here, and why
|
| 57 |
+
|
| 58 |
+
Training orchestration, data pipelines, the §12 memory-organ program (labeled ongoing
|
| 59 |
+
in the paper), and downstream control/steering machinery. The claims those support are
|
| 60 |
+
either reported with their own dated work-log provenance (paper, Appendix R) or not
|
| 61 |
+
yet published. This bundle is scoped to verify what the preprint asserts about these
|
| 62 |
+
frozen artifacts.
|
| 63 |
+
|
| 64 |
+
## Checksums / provenance
|
| 65 |
+
|
| 66 |
+
Both checkpoints are weights-only exports (optimizer state stripped) of the exact
|
| 67 |
+
training checkpoints named in the paper's Appendix R. Verify with:
|
| 68 |
+
`sha256sum -c SHA256SUMS`
|
SHA256SUMS
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
ddd042a6b8bb1a1e07cdebb75fad135e12955528e872f22b40e7d3fe15559535 release/cl33_oplm_chat_236m.pt
|
| 2 |
+
fe407328903731ecf94e5c7173a754311212a1c24ed2cbb92a1978faea6a1fbf release/cl33_oplm_prose_236m.pt
|
| 3 |
+
9120134a29141485043107b810ec08a01fb374bfe962421e45d9db88c0ca79ef release/invert_probe_chat.pt
|
| 4 |
+
df6030348482e9bf0e5ba37e2e62dfada8cf6550a0107a6ce86d68c81019bd8c model.py
|
| 5 |
+
69992c25815d9c77e11b5e8e7fee644c170387c630c3efe379e700c4d2a7000a model_v2.py
|
| 6 |
+
fa0a11eb90a78b8b36da5ae9ef2c92f92c99df0e0fe7e76fca9b8036d442bc9e repro_bottleneck.py
|
| 7 |
+
e93fcaaa8eae49a5bf401c8b722f3290d99f7142b3cab528cd06b02c39e9757f repro_reverse_readout.py
|
| 8 |
+
ff24dcfa83fb633e71fdeb0093c63b05d6118d574a1efd25fa5744978be0cc2f so33.py
|
| 9 |
+
fb8d17474cf20ec177a9d6c5810f8d66a79a4a0b7070fa4f111174d91d3b2d29 t3v3_wedge_memory.py
|
| 10 |
+
fa70959ec1657c7747b02ec78652c3ef445f7c72f661b36ea68278c4ab466147 tape_memory.py
|
| 11 |
+
692cd36f59622e37ea83cb7c2957928aa2d49dbf27bb4c362815e47cac96dc33 wedge.py
|
cl33_oplm_chat_236m.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:ddd042a6b8bb1a1e07cdebb75fad135e12955528e872f22b40e7d3fe15559535
|
| 3 |
+
size 946034587
|
cl33_oplm_prose_236m.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:fe407328903731ecf94e5c7173a754311212a1c24ed2cbb92a1978faea6a1fbf
|
| 3 |
+
size 945975203
|
invert_probe_chat.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:9120134a29141485043107b810ec08a01fb374bfe962421e45d9db88c0ca79ef
|
| 3 |
+
size 211878933
|
model.py
ADDED
|
@@ -0,0 +1,165 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Operator-emitting language model (cl33-opLM v0).
|
| 2 |
+
|
| 3 |
+
The OPERATOR-ONLY thesis, made structural: context reaches the next-token
|
| 4 |
+
prediction ONLY through emitted operators. A causal transformer "emitter"
|
| 5 |
+
produces, per position, a per-block so(3,3) bivector; those generate rotors; a
|
| 6 |
+
reversible matrix-action scan evolves a multi-block SO(3,3) state; the readout
|
| 7 |
+
sees ONLY that state. No residual bypass from the emitter to the readout — so
|
| 8 |
+
if the operators don't carry the information, PPL suffers (that IS the thesis).
|
| 9 |
+
|
| 10 |
+
tokens → embed → causal emitter → h_t
|
| 11 |
+
h_t → op_head → coefs_t (n_blocks, 15) [clipped, small init]
|
| 12 |
+
R_t = matrix_exp(Σ coefs_t·G) [per block, fp32]
|
| 13 |
+
s_t = R_t · s_{t-1} [reversible scan; s_0 learned]
|
| 14 |
+
logits = readout(LayerNorm(flatten s_t)) [readout sees ONLY the state]
|
| 15 |
+
|
| 16 |
+
Scan runs in float32 (bf16 diverges — proven on the lattice). Emitter may be
|
| 17 |
+
autocast-bf16; the state path is force-fp32.
|
| 18 |
+
"""
|
| 19 |
+
from __future__ import annotations
|
| 20 |
+
|
| 21 |
+
import math
|
| 22 |
+
import sys
|
| 23 |
+
from dataclasses import dataclass
|
| 24 |
+
from pathlib import Path
|
| 25 |
+
|
| 26 |
+
import torch
|
| 27 |
+
import torch.nn as nn
|
| 28 |
+
import torch.nn.functional as F
|
| 29 |
+
|
| 30 |
+
sys.path.insert(0, str(Path(__file__).resolve().parent))
|
| 31 |
+
from so33 import build_generators, rotor_from_coefs, N_GEN # type: ignore
|
| 32 |
+
from wedge import GradeTower # type: ignore
|
| 33 |
+
|
| 34 |
+
# per-block readout feature dims: g1 + g2 + g3 + g4 + g5 + g6
|
| 35 |
+
_GRADE1, _GRADE2, _GRADE3, _GRADE4, _GRADE5, _GRADE6 = 6, 15, 20, 15, 6, 1
|
| 36 |
+
_TOWER_DIM = _GRADE1 + _GRADE2 + _GRADE3 + _GRADE4 + _GRADE5 + _GRADE6 # 63
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
@dataclass
|
| 40 |
+
class OpEmitConfig:
|
| 41 |
+
vocab_size: int = 8192
|
| 42 |
+
d_model: int = 384
|
| 43 |
+
n_layers: int = 6
|
| 44 |
+
n_heads: int = 6
|
| 45 |
+
d_ff: int = 1536
|
| 46 |
+
max_seq_len: int = 256
|
| 47 |
+
n_blocks: int = 16 # multi-block SO(3,3) state (6·n_blocks dims)
|
| 48 |
+
coef_clip: float = 1.0 # per-block bivector L2-norm cap (bounds growth)
|
| 49 |
+
op_init_scale: float = 0.02 # rotors start ≈ identity
|
| 50 |
+
use_grade_tower: bool = True # readout sees full grade tower per block
|
| 51 |
+
dropout: float = 0.0
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
class EmitterBlock(nn.Module):
|
| 55 |
+
def __init__(self, c: OpEmitConfig):
|
| 56 |
+
super().__init__()
|
| 57 |
+
self.n_heads = c.n_heads
|
| 58 |
+
self.d_head = c.d_model // c.n_heads
|
| 59 |
+
self.qkv = nn.Linear(c.d_model, 3 * c.d_model, bias=False)
|
| 60 |
+
self.o = nn.Linear(c.d_model, c.d_model, bias=False)
|
| 61 |
+
self.norm1 = nn.LayerNorm(c.d_model)
|
| 62 |
+
self.norm2 = nn.LayerNorm(c.d_model)
|
| 63 |
+
self.mlp = nn.Sequential(nn.Linear(c.d_model, c.d_ff), nn.GELU(),
|
| 64 |
+
nn.Linear(c.d_ff, c.d_model))
|
| 65 |
+
|
| 66 |
+
def forward(self, x):
|
| 67 |
+
B, T, D = x.shape
|
| 68 |
+
h = self.norm1(x)
|
| 69 |
+
qkv = self.qkv(h).reshape(B, T, 3, self.n_heads, self.d_head)
|
| 70 |
+
q, k, v = qkv.unbind(2)
|
| 71 |
+
q, k, v = (t.transpose(1, 2) for t in (q, k, v)) # (B,H,T,dh)
|
| 72 |
+
a = F.scaled_dot_product_attention(q, k, v, is_causal=True)
|
| 73 |
+
a = a.transpose(1, 2).reshape(B, T, D)
|
| 74 |
+
x = x + self.o(a)
|
| 75 |
+
x = x + self.mlp(self.norm2(x))
|
| 76 |
+
return x
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
class OpEmitLM(nn.Module):
|
| 80 |
+
def __init__(self, c: OpEmitConfig):
|
| 81 |
+
super().__init__()
|
| 82 |
+
self.c = c
|
| 83 |
+
self.tok_embed = nn.Embedding(c.vocab_size, c.d_model)
|
| 84 |
+
self.pos_embed = nn.Embedding(c.max_seq_len, c.d_model)
|
| 85 |
+
self.blocks = nn.ModuleList([EmitterBlock(c) for _ in range(c.n_layers)])
|
| 86 |
+
self.norm = nn.LayerNorm(c.d_model)
|
| 87 |
+
self.op_head = nn.Linear(c.d_model, c.n_blocks * N_GEN)
|
| 88 |
+
nn.init.normal_(self.op_head.weight, std=1e-3)
|
| 89 |
+
nn.init.zeros_(self.op_head.bias)
|
| 90 |
+
# learned initial state per block (nonzero so rotors have something to act on)
|
| 91 |
+
self.s0 = nn.Parameter(torch.randn(c.n_blocks, 6) * 0.5)
|
| 92 |
+
feat_per_block = _TOWER_DIM if c.use_grade_tower else 6
|
| 93 |
+
self.feat_dim = c.n_blocks * feat_per_block
|
| 94 |
+
self.tower = GradeTower() if c.use_grade_tower else None
|
| 95 |
+
self.state_norm = nn.LayerNorm(self.feat_dim)
|
| 96 |
+
self.readout = nn.Linear(self.feat_dim, c.vocab_size, bias=False)
|
| 97 |
+
self.register_buffer("G", build_generators(dtype=torch.float32), persistent=False)
|
| 98 |
+
|
| 99 |
+
def emit(self, idx):
|
| 100 |
+
B, T = idx.shape
|
| 101 |
+
pos = torch.arange(T, device=idx.device)
|
| 102 |
+
x = self.tok_embed(idx) + self.pos_embed(pos)[None]
|
| 103 |
+
for blk in self.blocks:
|
| 104 |
+
x = blk(x)
|
| 105 |
+
h = self.norm(x)
|
| 106 |
+
coefs = self.op_head(h).view(B, T, self.c.n_blocks, N_GEN)
|
| 107 |
+
coefs = coefs * self.c.op_init_scale
|
| 108 |
+
# per-block L2-norm clip → bounded rotors → bounded state growth
|
| 109 |
+
n = coefs.norm(dim=-1, keepdim=True)
|
| 110 |
+
coefs = coefs * (self.c.coef_clip / n.clamp_min(self.c.coef_clip))
|
| 111 |
+
return coefs # (B,T,nb,15)
|
| 112 |
+
|
| 113 |
+
def scan(self, coefs, ablate=False):
|
| 114 |
+
"""Reversible matrix-action scan. Returns S (B,T,nb,6). fp32 forced."""
|
| 115 |
+
B, T, nb, _ = coefs.shape
|
| 116 |
+
coefs = coefs.float()
|
| 117 |
+
if ablate:
|
| 118 |
+
R = torch.eye(6, device=coefs.device).expand(B, T, nb, 6, 6)
|
| 119 |
+
else:
|
| 120 |
+
R = rotor_from_coefs(coefs, self.G.float()) # (B,T,nb,6,6)
|
| 121 |
+
s = self.s0.float().expand(B, nb, 6).contiguous()
|
| 122 |
+
s = s / s.norm(dim=-1, keepdim=True).clamp_min(1e-6)
|
| 123 |
+
S = torch.empty(B, T, nb, 6, device=coefs.device, dtype=torch.float32)
|
| 124 |
+
for t in range(T):
|
| 125 |
+
s = torch.einsum("bnij,bnj->bni", R[:, t], s)
|
| 126 |
+
# per-block unit-norm projection: bounds growth (boosts amplify ‖s‖),
|
| 127 |
+
# state lives on (S⁵)^nb. Matches the lattice's normalized recurrence.
|
| 128 |
+
s = s / s.norm(dim=-1, keepdim=True).clamp_min(1e-6)
|
| 129 |
+
S[:, t] = s
|
| 130 |
+
return S
|
| 131 |
+
|
| 132 |
+
def forward(self, idx, targets=None, ablate_scan=False):
|
| 133 |
+
"""ablate_scan=True → R→IDENTITY (state frozen at s0) but coefs/tower
|
| 134 |
+
stay LIVE. This is the BYPASS TEST: if the model still predicts well
|
| 135 |
+
with the rotors off, the grade tower is feeding operators to the readout
|
| 136 |
+
past the state recurrence (a bypass) — the causal claim is dead. Correct
|
| 137 |
+
pass: ablated PPL ≈ unigram floor. (Reported as absolute PPL, not a
|
| 138 |
+
ratio — the ratio inflates mechanically as trained PPL drops.)"""
|
| 139 |
+
B, T = idx.shape
|
| 140 |
+
coefs = self.emit(idx) # (B,T,nb,15)
|
| 141 |
+
S = self.scan(coefs, ablate=ablate_scan) # (B,T,nb,6) grade-1
|
| 142 |
+
if self.tower is not None:
|
| 143 |
+
# causal shifts of the operator sequence (grade-2)
|
| 144 |
+
B_prev = torch.zeros_like(coefs)
|
| 145 |
+
B_prev[:, 1:] = coefs[:, :-1]
|
| 146 |
+
B_prev2 = torch.zeros_like(coefs)
|
| 147 |
+
B_prev2[:, 2:] = coefs[:, :-2]
|
| 148 |
+
g3, g4, g5, g6 = self.tower(S, coefs, B_prev, B_prev2)
|
| 149 |
+
feat = torch.cat([S, coefs, g3, g4, g5, g6], dim=-1) # (B,T,nb,63)
|
| 150 |
+
feat = feat.reshape(B, T, -1)
|
| 151 |
+
else:
|
| 152 |
+
feat = S.reshape(B, T, -1)
|
| 153 |
+
feat = self.state_norm(feat)
|
| 154 |
+
logits = self.readout(feat) # (B,T,vocab)
|
| 155 |
+
loss = None
|
| 156 |
+
if targets is not None:
|
| 157 |
+
loss = F.cross_entropy(logits.reshape(-1, logits.shape[-1]),
|
| 158 |
+
targets.reshape(-1))
|
| 159 |
+
return logits, loss
|
| 160 |
+
|
| 161 |
+
def param_count(self):
|
| 162 |
+
return sum(p.numel() for p in self.parameters() if p.requires_grad)
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
__all__ = ["OpEmitConfig", "OpEmitLM"]
|
model_v2.py
ADDED
|
@@ -0,0 +1,332 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""cl33-opLM v2 — operator-native attention over the state trajectory.
|
| 2 |
+
|
| 3 |
+
The v0/widen finding: the operator-only bottleneck forces a RECURRENT state that
|
| 4 |
+
compresses history; widening independent blocks doesn't beat compression. The
|
| 5 |
+
fix (Garret): restore attention, but map it 1:1 onto the algebra so it stays
|
| 6 |
+
operator-derived (causal transparency preserved).
|
| 7 |
+
|
| 8 |
+
Operator-native attention (per block = per head):
|
| 9 |
+
state scan: s_t = R(B_state_t) · s_{t-1} (as before, reversible)
|
| 10 |
+
query/key: Q_t = R(B_q_t) · s_t, K_s = R(B_k_s) · s_s (emitted rotors)
|
| 11 |
+
score: ⟨Q_t, K_s⟩_η (η-metric inner product, causal)
|
| 12 |
+
attend: a_t = Σ_s softmax(score)_ts · s_s (values = the states)
|
| 13 |
+
readout: LayerNorm(flatten a_t) → vocab
|
| 14 |
+
|
| 15 |
+
Because rotors preserve η, score(t,s) = ⟨s_t, (R_q⁻¹R_k)·s_s⟩_η — attention is
|
| 16 |
+
the alignment of the current state with a LEARNED-RELATIVELY-ROTATED past state.
|
| 17 |
+
Everything is operator-derived → operator-only bottleneck holds → the R_state→
|
| 18 |
+
identity bypass test still collapses to unigram (all states=s0 ⇒ attention over
|
| 19 |
+
identical values ⇒ constant readout).
|
| 20 |
+
|
| 21 |
+
Scan + attention in fp32 (bf16 proven-unsafe on the recurrence).
|
| 22 |
+
"""
|
| 23 |
+
from __future__ import annotations
|
| 24 |
+
|
| 25 |
+
import sys
|
| 26 |
+
from dataclasses import dataclass
|
| 27 |
+
from pathlib import Path
|
| 28 |
+
|
| 29 |
+
import torch
|
| 30 |
+
import torch.nn as nn
|
| 31 |
+
import torch.nn.functional as F
|
| 32 |
+
|
| 33 |
+
sys.path.insert(0, str(Path(__file__).resolve().parent))
|
| 34 |
+
from so33 import build_generators, rotor_from_coefs, N_GEN, ETA # type: ignore
|
| 35 |
+
from model import EmitterBlock # reuse the emitter block
|
| 36 |
+
from wedge import GradeTower # grade tower for short-context (current operator)
|
| 37 |
+
from t3v3_wedge_memory import WedgeMemory, TokenCopyMemory # associative / copy memory
|
| 38 |
+
from tape_memory import TapeMemory # reversible-tape read (address by token, flow algebra)
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
@dataclass
|
| 42 |
+
class OpEmitV2Config:
|
| 43 |
+
vocab_size: int = 8192
|
| 44 |
+
d_model: int = 384
|
| 45 |
+
n_layers: int = 6
|
| 46 |
+
n_heads: int = 6
|
| 47 |
+
d_ff: int = 1536
|
| 48 |
+
max_seq_len: int = 256
|
| 49 |
+
n_blocks: int = 32 # = attention heads (each attends in its SO(3,3))
|
| 50 |
+
coef_clip: float = 1.0
|
| 51 |
+
op_init_scale: float = 0.02
|
| 52 |
+
qk_init_scale: float = 0.1 # q/k rotors can be larger (learn what to attend)
|
| 53 |
+
use_grade_tower: bool = True # v2.2: current-operator grade tower in readout
|
| 54 |
+
use_wedge_memory: bool = False # wedge bivector associative memory (KV binding)
|
| 55 |
+
wedge_key_source: str = "operator" # "operator" | "state" | "token"
|
| 56 |
+
wedge_value_source: str = "same" # "same" | "operator" (G1: algebra value w/ token key)
|
| 57 |
+
wedge_delta_rule: bool = False # DeltaNet residual write (D3 interference fix)
|
| 58 |
+
use_token_copy: bool = False # token-content copy channel (measured transparency cost)
|
| 59 |
+
use_tape_memory: bool = False # reversible-tape read (REVERSIBLE_TAPE_DESIGN.md)
|
| 60 |
+
tape_value_mode: str = "increment" # "displacement" | "state" | "increment"
|
| 61 |
+
tape_dual_address: bool = False # v2.5: + token-faithful exact-match channel (TAPE_ADDRESSING_V25.md)
|
| 62 |
+
tape_addr: str = "token" # v2.8: "token" (address by token embedding) | "operator" (address
|
| 63 |
+
# by the emitted LM operator — native relational key, not the
|
| 64 |
+
# redundant token channel CE routes around). Pairs with tape_compose.
|
| 65 |
+
tape_compose: bool = False # v2.8: gp-COMPOSE recalled operator with the query rotor before
|
| 66 |
+
# readout (genesis marriage compose + v2.7 product-with-query). The
|
| 67 |
+
# recalled op stops being read in isolation — it composes with the
|
| 68 |
+
# current computation. Needs tape_value_mode=increment (15-d recall).
|
| 69 |
+
tape_compose_mode: str = "add" # "add" = additive zero-init logit (SAFE, out of LN; augments the
|
| 70 |
+
# distribution). "tilt" = recalled operator (gated, zero-init) TILTS
|
| 71 |
+
# the decoded state itself, decoded by the normal readout (augments
|
| 72 |
+
# BEHAVIOR; in the LN → tests whether it survives tape-silencing).
|
| 73 |
+
scan_only: bool = False # zero the O(T²) attention — tape/scan carry history
|
| 74 |
+
dropout: float = 0.0
|
| 75 |
+
grad_checkpoint: bool = False # activation-checkpoint the emitter transformer blocks: recompute
|
| 76 |
+
# in backward instead of storing (bit-identical forward, unchanged
|
| 77 |
+
# architecture/reversibility). Frees the dominant activation memory
|
| 78 |
+
# → larger batch → better sequential-scan SM occupancy on small GPUs.
|
| 79 |
+
|
| 80 |
+
|
| 81 |
+
class OpEmitLMv2(nn.Module):
|
| 82 |
+
def __init__(self, c: OpEmitV2Config):
|
| 83 |
+
super().__init__()
|
| 84 |
+
self.c = c
|
| 85 |
+
self.tok_embed = nn.Embedding(c.vocab_size, c.d_model)
|
| 86 |
+
self.pos_embed = nn.Embedding(c.max_seq_len, c.d_model)
|
| 87 |
+
self.blocks = nn.ModuleList([EmitterBlock(c) for _ in range(c.n_layers)])
|
| 88 |
+
self.norm = nn.LayerNorm(c.d_model)
|
| 89 |
+
# emit 3 bivectors per block: state-evolution, query-rotor, key-rotor
|
| 90 |
+
self.op_head = nn.Linear(c.d_model, c.n_blocks * 3 * N_GEN)
|
| 91 |
+
nn.init.normal_(self.op_head.weight, std=1e-3)
|
| 92 |
+
nn.init.zeros_(self.op_head.bias)
|
| 93 |
+
self.s0 = nn.Parameter(torch.randn(c.n_blocks, 6) * 0.5)
|
| 94 |
+
# readout features per block:
|
| 95 |
+
# attended(6) + s_t(6=g1) [residual: long-context attn + current state]
|
| 96 |
+
# + grade tower g2(15)+g3(20)+g4(15)+g5(6)+g6(1)=57 [v2.2: current
|
| 97 |
+
# operator, rich at pos 1 — fixes the rotation-of-s0 short-context
|
| 98 |
+
# limit]. All operator-derived → bottleneck preserved.
|
| 99 |
+
self.tower = GradeTower() if c.use_grade_tower else None
|
| 100 |
+
self.wedge = (WedgeMemory(key_source=getattr(c, "wedge_key_source", "operator"),
|
| 101 |
+
value_source=getattr(c, "wedge_value_source", "same"),
|
| 102 |
+
delta_rule=getattr(c, "wedge_delta_rule", False),
|
| 103 |
+
d_model=c.d_model, nb=c.n_blocks)
|
| 104 |
+
if getattr(c, "use_wedge_memory", False) else None)
|
| 105 |
+
self.token_copy = (TokenCopyMemory(c.d_model, c.n_blocks)
|
| 106 |
+
if getattr(c, "use_token_copy", False) else None)
|
| 107 |
+
self.tape = (TapeMemory(c.d_model, c.n_blocks,
|
| 108 |
+
value_mode=getattr(c, "tape_value_mode", "increment"),
|
| 109 |
+
dual_address=getattr(c, "tape_dual_address", False),
|
| 110 |
+
addr=getattr(c, "tape_addr", "token"))
|
| 111 |
+
if getattr(c, "use_tape_memory", False) else None)
|
| 112 |
+
# readout features per block: +6 each for the wedge read and the copy read.
|
| 113 |
+
# The TAPE is deliberately NOT in this vector — it is a SEPARATE post-LayerNorm additive
|
| 114 |
+
# readout term (see assemble). Reason: concatenating the tape into the shared LayerNorm lets
|
| 115 |
+
# any tape contribution perturb the normalization of the base features, so the optimizer
|
| 116 |
+
# silences the tape to protect the base LM (observed: v2.4 gate 0.01→0.0005). Keeping the
|
| 117 |
+
# tape OUT of the LN leaves the base bit-exact (clean warm-start) and lets the tape co-train.
|
| 118 |
+
feat_per_block = (12 + (6 if self.wedge is not None else 0)
|
| 119 |
+
+ (6 if self.token_copy is not None else 0)
|
| 120 |
+
+ (57 if c.use_grade_tower else 0))
|
| 121 |
+
self.state_norm = nn.LayerNorm(c.n_blocks * feat_per_block)
|
| 122 |
+
self.readout = nn.Linear(c.n_blocks * feat_per_block, c.vocab_size, bias=False)
|
| 123 |
+
# tape = additive side channel, ZERO-INIT readout → off at step 0, learns on. The zero-init
|
| 124 |
+
# readout IS the clean off-switch (no gate scalar needed, no LN-suppression dynamic).
|
| 125 |
+
if self.tape is not None:
|
| 126 |
+
self.tape_readout = nn.Linear(self.tape.out_dim() * c.n_blocks, c.vocab_size, bias=False)
|
| 127 |
+
nn.init.zeros_(self.tape_readout.weight)
|
| 128 |
+
# v2.8 compose channel: recalled operator gp-composed with the query rotor, applied to state.
|
| 129 |
+
# Separate ZERO-INIT readout (same off-switch pattern) → step-0 ≡ base, co-trains on.
|
| 130 |
+
self.tape_compose = getattr(c, "tape_compose", False) and self.tape is not None
|
| 131 |
+
if self.tape_compose:
|
| 132 |
+
assert self.tape.out_dim() == 15, "tape_compose needs 15-d bivector recall (tape_value_mode=increment)"
|
| 133 |
+
if getattr(c, "tape_compose_mode", "add") == "tilt":
|
| 134 |
+
# gated recalled operator tilts the decoded state; per-block gate ZERO-INIT → R_tilt=I,
|
| 135 |
+
# step-0 state unchanged. In the LN via parts=[a, S_tilt] → the behavior-augmenting arm.
|
| 136 |
+
self.tape_tilt_gate = nn.Parameter(torch.zeros(c.n_blocks, 1))
|
| 137 |
+
else:
|
| 138 |
+
self.tape_compose_readout = nn.Linear(c.n_blocks * 6, c.vocab_size, bias=False)
|
| 139 |
+
nn.init.zeros_(self.tape_compose_readout.weight)
|
| 140 |
+
self.register_buffer("G", build_generators(dtype=torch.float32), persistent=False)
|
| 141 |
+
self.register_buffer("eta", ETA.clone(), persistent=False)
|
| 142 |
+
|
| 143 |
+
def emit(self, idx):
|
| 144 |
+
B, T = idx.shape
|
| 145 |
+
pos = torch.arange(T, device=idx.device)
|
| 146 |
+
x = self.tok_embed(idx) + self.pos_embed(pos)[None]
|
| 147 |
+
if getattr(self.c, "grad_checkpoint", False) and self.training:
|
| 148 |
+
from torch.utils.checkpoint import checkpoint
|
| 149 |
+
for blk in self.blocks:
|
| 150 |
+
x = checkpoint(blk, x, use_reentrant=False) # recompute in backward, free activations
|
| 151 |
+
else:
|
| 152 |
+
for blk in self.blocks:
|
| 153 |
+
x = blk(x)
|
| 154 |
+
h = self.norm(x)
|
| 155 |
+
raw = self.op_head(h).view(B, T, self.c.n_blocks, 3, N_GEN)
|
| 156 |
+
Bs, Bq, Bk = raw.unbind(-2) # each (B,T,nb,15)
|
| 157 |
+
|
| 158 |
+
def clip(bv, scale):
|
| 159 |
+
bv = bv * scale
|
| 160 |
+
n = bv.norm(dim=-1, keepdim=True)
|
| 161 |
+
return bv * (self.c.coef_clip / n.clamp_min(self.c.coef_clip))
|
| 162 |
+
return (clip(Bs, self.c.op_init_scale),
|
| 163 |
+
clip(Bq, self.c.qk_init_scale),
|
| 164 |
+
clip(Bk, self.c.qk_init_scale))
|
| 165 |
+
|
| 166 |
+
@torch.no_grad()
|
| 167 |
+
def emitter_hidden(self, idx):
|
| 168 |
+
"""The emitter's final hidden h_t (baseline for recoverability probes)."""
|
| 169 |
+
B, T = idx.shape
|
| 170 |
+
pos = torch.arange(T, device=idx.device)
|
| 171 |
+
x = self.tok_embed(idx) + self.pos_embed(pos)[None]
|
| 172 |
+
for blk in self.blocks:
|
| 173 |
+
x = blk(x)
|
| 174 |
+
return self.norm(x) # (B,T,d_model)
|
| 175 |
+
|
| 176 |
+
def scan(self, Bs, ablate=False):
|
| 177 |
+
"""Reversible state recurrence. Returns S (B,T,nb,6), fp32."""
|
| 178 |
+
B, T, nb, _ = Bs.shape
|
| 179 |
+
Bs = Bs.float()
|
| 180 |
+
R = (torch.eye(6, device=Bs.device).expand(B, T, nb, 6, 6) if ablate
|
| 181 |
+
else rotor_from_coefs(Bs, self.G.float()))
|
| 182 |
+
s = self.s0.float().expand(B, nb, 6).contiguous()
|
| 183 |
+
s = s / s.norm(dim=-1, keepdim=True).clamp_min(1e-6)
|
| 184 |
+
S = torch.empty(B, T, nb, 6, device=Bs.device, dtype=torch.float32)
|
| 185 |
+
for t in range(T):
|
| 186 |
+
s = torch.einsum("bnij,bnj->bni", R[:, t], s)
|
| 187 |
+
s = s / s.norm(dim=-1, keepdim=True).clamp_min(1e-6)
|
| 188 |
+
S[:, t] = s
|
| 189 |
+
return S
|
| 190 |
+
|
| 191 |
+
def scan_parallel(self, Bs, ablate=False):
|
| 192 |
+
"""Parallel associative scan (PARALLEL_SCAN_SPEC): the per-step normalize cancels, so
|
| 193 |
+
S_t = normalize(P_t·s0) with P_t = R_t···R_0 a prefix product. Hillis-Steele scan under the
|
| 194 |
+
associative operator A∘B = normalize_F(A·B) (Frobenius-normalized to avoid boost overflow).
|
| 195 |
+
O(log T) depth. GATED against scan() — must match bit-for-bit before it replaces it."""
|
| 196 |
+
B, T, nb, _ = Bs.shape
|
| 197 |
+
Bs = Bs.float()
|
| 198 |
+
R = (torch.eye(6, device=Bs.device).expand(B, T, nb, 6, 6).contiguous() if ablate
|
| 199 |
+
else rotor_from_coefs(Bs, self.G.float()))
|
| 200 |
+
def nF(M):
|
| 201 |
+
return M / (M.reshape(*M.shape[:-2], 36).norm(dim=-1)[..., None, None] + 1e-30)
|
| 202 |
+
P = nF(R) # (B,T,nb,6,6)
|
| 203 |
+
I = torch.eye(6, device=Bs.device, dtype=P.dtype).expand(B, 1, nb, 6, 6)
|
| 204 |
+
idx = torch.arange(T, device=Bs.device)
|
| 205 |
+
d = 1
|
| 206 |
+
while d < T:
|
| 207 |
+
right = torch.cat([I.expand(B, d, nb, 6, 6), P[:, :T - d]], dim=1) # right[t]=P[t-d], I for t<d
|
| 208 |
+
combined = nF(P @ right) # newest(left) @ older(right), correct order
|
| 209 |
+
mask = (idx >= d)[None, :, None, None, None]
|
| 210 |
+
P = torch.where(mask, combined, P)
|
| 211 |
+
d *= 2
|
| 212 |
+
s0 = self.s0.float().expand(B, nb, 6)
|
| 213 |
+
s0 = s0 / s0.norm(dim=-1, keepdim=True).clamp_min(1e-6)
|
| 214 |
+
S = torch.einsum("btnij,bnj->btni", P, s0)
|
| 215 |
+
return S / S.norm(dim=-1, keepdim=True).clamp_min(1e-6)
|
| 216 |
+
|
| 217 |
+
def attend(self, S, Bq, Bk):
|
| 218 |
+
"""Operator-native causal attention over the state trajectory.
|
| 219 |
+
S (B,T,nb,6); Bq,Bk (B,T,nb,15). Returns attended (B,T,nb,6)."""
|
| 220 |
+
B, T, nb, _ = S.shape
|
| 221 |
+
Rq = rotor_from_coefs(Bq.float(), self.G.float()) # (B,T,nb,6,6)
|
| 222 |
+
Rk = rotor_from_coefs(Bk.float(), self.G.float())
|
| 223 |
+
Q = torch.einsum("btnij,btnj->btni", Rq, S) # (B,T,nb,6)
|
| 224 |
+
K = torch.einsum("btnij,btnj->btni", Rk, S)
|
| 225 |
+
# η-metric scores: ⟨Q_t, K_s⟩_η, per block/head. (B,nb,T,T)
|
| 226 |
+
Kw = K * self.eta.to(K.dtype) # apply η to keys
|
| 227 |
+
scores = torch.einsum("btni,bsni->bnts", Q, Kw) / (6 ** 0.5)
|
| 228 |
+
causal = torch.triu(torch.ones(T, T, device=S.device, dtype=torch.bool), 1)
|
| 229 |
+
scores = scores.masked_fill(causal, float("-inf"))
|
| 230 |
+
A = F.softmax(scores, dim=-1) # (B,nb,T,T)
|
| 231 |
+
attended = torch.einsum("bnts,bsni->btni", A, S) # values = states
|
| 232 |
+
return attended
|
| 233 |
+
|
| 234 |
+
def assemble(self, Bs, Bq, Bk, scan_only=False, tok_emb=None):
|
| 235 |
+
"""Full pass from emitted operators → logits (scan + attention + tower
|
| 236 |
+
+ readout). Separated from emit() so CONTROL interventions can perturb
|
| 237 |
+
the operators and re-run only the downstream. Returns (B,T,vocab).
|
| 238 |
+
|
| 239 |
+
scan_only=True zeroes the cross-position attention path, so recall must
|
| 240 |
+
come from the recurrent reversible state alone (the fair vs-xLSTM memory
|
| 241 |
+
test — isolates the linear-recurrent memory from the O(T²) attention)."""
|
| 242 |
+
B, T = Bs.shape[:2]
|
| 243 |
+
S = self.scan(Bs, ablate=False) # (B,T,nb,6)
|
| 244 |
+
# tape read (computed here so a 'tilt' compose can act on the decoded state below)
|
| 245 |
+
tape_read = None
|
| 246 |
+
if self.tape is not None and tok_emb is not None:
|
| 247 |
+
R_state = (rotor_from_coefs(Bs.float(), self.G.float())
|
| 248 |
+
if (self.tape.value_mode in ("displacement", "multi")
|
| 249 |
+
or getattr(self.tape, "addr", "token") == "target") else None)
|
| 250 |
+
tape_read = self.tape(S, Bs, R_state, tok_emb, self.eta, Bq=Bq) # (B,T,nb,out_dim)
|
| 251 |
+
# v2.8 tilt-compose: the recalled operator (gated, zero-init) tilts the state that gets decoded
|
| 252 |
+
# by the NORMAL readout — augments behavior, then decoded like usual.
|
| 253 |
+
S_ro = S
|
| 254 |
+
if (tape_read is not None and getattr(self, "tape_compose", False)
|
| 255 |
+
and getattr(self.c, "tape_compose_mode", "add") == "tilt"):
|
| 256 |
+
tilt_biv = self.tape_tilt_gate * tape_read.float() # (B,T,nb,15); gate 0 → identity
|
| 257 |
+
R_tilt = rotor_from_coefs(tilt_biv, self.G.float()) # (B,T,nb,6,6)
|
| 258 |
+
S_ro = torch.einsum("btnij,btnj->btni", R_tilt, S) # recalled op tilts the decoded state
|
| 259 |
+
a = torch.zeros_like(S) if scan_only else self.attend(S, Bq, Bk)
|
| 260 |
+
parts = [a, S_ro]
|
| 261 |
+
if self.wedge is not None:
|
| 262 |
+
# O(T) linear-recurrent associative memory. operator mode keys on the
|
| 263 |
+
# per-token emitted operators (content-bearing); state mode (ablation)
|
| 264 |
+
# keys on the transported state + adjoint-transports the memory.
|
| 265 |
+
Rq = Rk = R_state = None
|
| 266 |
+
if self.wedge.key_source == "state":
|
| 267 |
+
G32 = self.G.float()
|
| 268 |
+
Rq = rotor_from_coefs(Bq.float(), G32)
|
| 269 |
+
Rk = rotor_from_coefs(Bk.float(), G32)
|
| 270 |
+
R_state = rotor_from_coefs(Bs.float(), G32)
|
| 271 |
+
parts.append(self.wedge(S, Rq, Rk, self.eta, R_state=R_state,
|
| 272 |
+
Bs=Bs, Bq=Bq, Bk=Bk, tok_emb=tok_emb)) # r_t (B,T,nb,6)
|
| 273 |
+
if self.token_copy is not None and tok_emb is not None:
|
| 274 |
+
# token-content copy channel (bypasses operators; transparency cost measured)
|
| 275 |
+
parts.append(self.token_copy(tok_emb)) # (B,T,nb,6)
|
| 276 |
+
# (tape_read computed above, before the state-tilt; consumed as a SEPARATE post-LN additive
|
| 277 |
+
# term below in "add" mode, or already applied as a state tilt above in "tilt" mode.)
|
| 278 |
+
if self.tower is not None:
|
| 279 |
+
Bs_prev = torch.zeros_like(Bs); Bs_prev[:, 1:] = Bs[:, :-1]
|
| 280 |
+
Bs_prev2 = torch.zeros_like(Bs); Bs_prev2[:, 2:] = Bs[:, :-2]
|
| 281 |
+
g3, g4, g5, g6 = self.tower(S, Bs, Bs_prev, Bs_prev2)
|
| 282 |
+
parts += [Bs, g3, g4, g5, g6]
|
| 283 |
+
feat = self.state_norm(torch.cat(parts, dim=-1).reshape(B, T, -1))
|
| 284 |
+
logits = self.readout(feat)
|
| 285 |
+
if tape_read is not None: # additive, post-LN, zero-init
|
| 286 |
+
logits = logits + self.tape_readout(tape_read.reshape(B, T, -1).to(logits.dtype))
|
| 287 |
+
if (tape_read is not None and getattr(self, "tape_compose", False)
|
| 288 |
+
and getattr(self.c, "tape_compose_mode", "add") == "add"):
|
| 289 |
+
# gp-COMPOSE the recalled operator with the current query rotor (exact matrix/rotor
|
| 290 |
+
# composition — lossless in the grade-1 action), apply to the state. The recalled
|
| 291 |
+
# operation now composes with the current computation instead of being read in isolation:
|
| 292 |
+
# genesis-marriage compose + v2.7's product-with-query. Zero-init readout → off at step 0.
|
| 293 |
+
G32 = self.G.float()
|
| 294 |
+
R_rec = rotor_from_coefs(tape_read.float(), G32) # (B,T,nb,6,6) recalled operator
|
| 295 |
+
R_q = rotor_from_coefs(Bq.float(), G32) # (B,T,nb,6,6) current query
|
| 296 |
+
R_comp = torch.einsum("btnij,btnjk->btnik", R_rec, R_q) # exact rotor composition
|
| 297 |
+
comp = torch.einsum("btnij,btnj->btni", R_comp, S) # composed op applied to state
|
| 298 |
+
logits = logits + self.tape_compose_readout(comp.reshape(B, T, -1).to(logits.dtype))
|
| 299 |
+
return logits
|
| 300 |
+
|
| 301 |
+
def forward(self, idx, targets=None, ablate_scan=False, amp_emit=False):
|
| 302 |
+
# amp_emit: run the transformer emitter in fp16 (tensor-core speedup, safe)
|
| 303 |
+
# but the operator algebra (scan/attend/tower/readout) in fp32 — fp16 there
|
| 304 |
+
# overflows (η-metric scores + nested wedge products) → NaN. Split keeps the
|
| 305 |
+
# bulk of compute fast while the delicate algebra stays numerically exact.
|
| 306 |
+
if amp_emit:
|
| 307 |
+
with torch.autocast("cuda", dtype=torch.float16):
|
| 308 |
+
Bs, Bq, Bk = self.emit(idx)
|
| 309 |
+
Bs, Bq, Bk = Bs.float(), Bq.float(), Bk.float()
|
| 310 |
+
else:
|
| 311 |
+
Bs, Bq, Bk = self.emit(idx)
|
| 312 |
+
if ablate_scan:
|
| 313 |
+
Bs = torch.zeros_like(Bs) # transparency test
|
| 314 |
+
# tok_emb (raw embedding) feeds the token-copy / token-key channels; on the
|
| 315 |
+
# ablate path it is retained → bypass ratio then reflects how much those
|
| 316 |
+
# channels carry prediction around the (zeroed) operators.
|
| 317 |
+
tok_emb = self.tok_embed(idx) if (self.token_copy is not None
|
| 318 |
+
or self.tape is not None
|
| 319 |
+
or (self.wedge is not None and self.wedge.key_source == "token")) else None
|
| 320 |
+
logits = self.assemble(Bs, Bq, Bk, scan_only=getattr(self.c, "scan_only", False),
|
| 321 |
+
tok_emb=tok_emb)
|
| 322 |
+
loss = None
|
| 323 |
+
if targets is not None:
|
| 324 |
+
loss = F.cross_entropy(logits.reshape(-1, logits.shape[-1]),
|
| 325 |
+
targets.reshape(-1))
|
| 326 |
+
return logits, loss
|
| 327 |
+
|
| 328 |
+
def param_count(self):
|
| 329 |
+
return sum(p.numel() for p in self.parameters() if p.requires_grad)
|
| 330 |
+
|
| 331 |
+
|
| 332 |
+
__all__ = ["OpEmitV2Config", "OpEmitLMv2"]
|
repro_bottleneck.py
ADDED
|
@@ -0,0 +1,59 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Reproduce the mandatory-operator-bottleneck claim on a released cl33-opLM checkpoint.
|
| 2 |
+
|
| 3 |
+
Zeroing the emitted operators (Bs=0 everywhere: identity rotors in the scan, zeroed
|
| 4 |
+
readout features) multiplies perplexity by orders of magnitude — the model has no
|
| 5 |
+
other path to output. Paper: "One Object" §1/§2/§7; 314x measured on the chat
|
| 6 |
+
checkpoint (prose validation); expect the same order of magnitude on WikiText-103.
|
| 7 |
+
|
| 8 |
+
Usage: python repro_bottleneck.py --ckpt cl33_oplm_chat_236m.pt [--iters 20]
|
| 9 |
+
Deps: torch, transformers, datasets (model_v2.py + so33.py from this bundle)
|
| 10 |
+
"""
|
| 11 |
+
import argparse, math, sys
|
| 12 |
+
from pathlib import Path
|
| 13 |
+
import torch, torch.nn.functional as F
|
| 14 |
+
sys.path.insert(0, str(Path(__file__).resolve().parent))
|
| 15 |
+
from model_v2 import OpEmitV2Config, OpEmitLMv2
|
| 16 |
+
|
| 17 |
+
ap = argparse.ArgumentParser()
|
| 18 |
+
ap.add_argument("--ckpt", default="cl33_oplm_chat_236m.pt")
|
| 19 |
+
ap.add_argument("--iters", type=int, default=20)
|
| 20 |
+
ap.add_argument("--seq", type=int, default=1024)
|
| 21 |
+
a = ap.parse_args()
|
| 22 |
+
dev = "cuda" if torch.cuda.is_available() else "cpu"
|
| 23 |
+
|
| 24 |
+
d = torch.load(a.ckpt, map_location="cpu", weights_only=False)
|
| 25 |
+
cfg = OpEmitV2Config(**{k: v for k, v in d["config"].items()
|
| 26 |
+
if k in OpEmitV2Config.__dataclass_fields__})
|
| 27 |
+
m = OpEmitLMv2(cfg); m.load_state_dict(d["model"]); m.eval().to(dev)
|
| 28 |
+
for p in m.parameters(): p.requires_grad_(False)
|
| 29 |
+
print(f"loaded {a.ckpt} | {sum(p.numel() for p in m.parameters())/1e6:.1f}M params | step {d.get('step')}")
|
| 30 |
+
|
| 31 |
+
from transformers import GPT2TokenizerFast
|
| 32 |
+
from datasets import load_dataset
|
| 33 |
+
tok = GPT2TokenizerFast.from_pretrained("gpt2")
|
| 34 |
+
text = "\n\n".join(load_dataset("Salesforce/wikitext", "wikitext-103-raw-v1",
|
| 35 |
+
split="test")["text"])
|
| 36 |
+
ids = tok(text, add_special_tokens=False)["input_ids"]
|
| 37 |
+
print(f"wikitext-103 test: {len(ids)/1e6:.2f}M tokens")
|
| 38 |
+
|
| 39 |
+
def ce_pass(zero_ops: bool):
|
| 40 |
+
tot, n = 0.0, 0
|
| 41 |
+
for i in range(a.iters):
|
| 42 |
+
s = i * a.seq
|
| 43 |
+
x = torch.tensor([ids[s:s+a.seq]], device=dev)
|
| 44 |
+
y = torch.tensor([ids[s+1:s+a.seq+1]], device=dev)
|
| 45 |
+
with torch.no_grad():
|
| 46 |
+
Bs, Bq, Bk = m.emit(x)
|
| 47 |
+
if zero_ops: Bs = torch.zeros_like(Bs)
|
| 48 |
+
logits = m.assemble(Bs, Bq, Bk, scan_only=getattr(cfg, "scan_only", False))
|
| 49 |
+
if isinstance(logits, tuple): logits = logits[0]
|
| 50 |
+
tot += float(F.cross_entropy(logits.reshape(-1, logits.shape[-1]).float(),
|
| 51 |
+
y.reshape(-1), reduction="sum"))
|
| 52 |
+
n += y.numel()
|
| 53 |
+
return tot / n
|
| 54 |
+
|
| 55 |
+
ce_nat = ce_pass(False); ce_off = ce_pass(True)
|
| 56 |
+
print(f"\n native : CE {ce_nat:.3f} PPL {math.exp(ce_nat):9.1f}")
|
| 57 |
+
print(f" ops off: CE {ce_off:.3f} PPL {math.exp(ce_off):9.1f}")
|
| 58 |
+
print(f" BOTTLENECK RATIO: {math.exp(ce_off)/math.exp(ce_nat):.0f}x "
|
| 59 |
+
f"(paper: 314x on chat ckpt / prose val; order of magnitude is the claim)")
|
repro_reverse_readout.py
ADDED
|
@@ -0,0 +1,47 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Reproduce the operator-transcript claim: decode text back out of the operator
|
| 2 |
+
stream alone. Paper §7 (inversion battery, 0.86 top-1/50k on scale-C; the released
|
| 3 |
+
probe is trained on the chat checkpoint's operators — held-out top-1 0.860).
|
| 4 |
+
|
| 5 |
+
Usage: python repro_reverse_readout.py [--text "any text you like"]
|
| 6 |
+
"""
|
| 7 |
+
import argparse, sys
|
| 8 |
+
from pathlib import Path
|
| 9 |
+
import torch
|
| 10 |
+
sys.path.insert(0, str(Path(__file__).resolve().parent))
|
| 11 |
+
from model_v2 import OpEmitV2Config, OpEmitLMv2
|
| 12 |
+
|
| 13 |
+
ap = argparse.ArgumentParser()
|
| 14 |
+
ap.add_argument("--ckpt", default="cl33_oplm_chat_236m.pt")
|
| 15 |
+
ap.add_argument("--probe", default="invert_probe_chat.pt")
|
| 16 |
+
ap.add_argument("--text", default="The reverse readout decodes the conversation "
|
| 17 |
+
"from the operator record alone, with no token input.")
|
| 18 |
+
a = ap.parse_args()
|
| 19 |
+
dev = "cuda" if torch.cuda.is_available() else "cpu"
|
| 20 |
+
|
| 21 |
+
d = torch.load(a.ckpt, map_location="cpu", weights_only=False)
|
| 22 |
+
cfg = OpEmitV2Config(**{k: v for k, v in d["config"].items()
|
| 23 |
+
if k in OpEmitV2Config.__dataclass_fields__})
|
| 24 |
+
m = OpEmitLMv2(cfg); m.load_state_dict(d["model"]); m.eval().to(dev)
|
| 25 |
+
import torch.nn as nn
|
| 26 |
+
pd = torch.load(a.probe, map_location="cpu", weights_only=False)
|
| 27 |
+
probe = nn.Sequential(nn.Linear(pd["in_dim"], pd["hidden"]), nn.GELU(),
|
| 28 |
+
nn.Linear(pd["hidden"], pd["vocab"]))
|
| 29 |
+
probe.load_state_dict(pd["probe"]); probe.eval().to(dev)
|
| 30 |
+
# stored held-out results ride with the artifact:
|
| 31 |
+
print("probe card:", {k: round(v, 3) for k, v in pd["results"].items()
|
| 32 |
+
if isinstance(v, float)})
|
| 33 |
+
|
| 34 |
+
from transformers import GPT2TokenizerFast
|
| 35 |
+
tok = GPT2TokenizerFast.from_pretrained("gpt2")
|
| 36 |
+
ids = tok(a.text, add_special_tokens=False)["input_ids"]
|
| 37 |
+
x = torch.tensor([ids], device=dev)
|
| 38 |
+
with torch.no_grad():
|
| 39 |
+
Bs, Bq, Bk = m.emit(x)
|
| 40 |
+
t = len(ids)
|
| 41 |
+
feats = torch.cat([Bs.reshape(1, t, -1), Bq.reshape(1, t, -1),
|
| 42 |
+
Bk.reshape(1, t, -1)], -1).float() # [Bs_480|Bq_480|Bk_480]
|
| 43 |
+
pred = probe(feats).argmax(-1)[0]
|
| 44 |
+
hits = sum(int(p == t) for p, t in zip(pred.tolist(), ids))
|
| 45 |
+
print(f"tokens: {len(ids)} | decoded from operators alone, top-1 exact: "
|
| 46 |
+
f"{hits}/{len(ids)} = {hits/len(ids):.2f}")
|
| 47 |
+
print("decoded:", tok.decode(pred.tolist()))
|
so33.py
ADDED
|
@@ -0,0 +1,66 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""so(3,3) generator basis + matrix-exp rotor + reversible matrix-action scan.
|
| 2 |
+
|
| 3 |
+
The operator-emitting LM works in a multi-block SO(3,3) state: each block is a
|
| 4 |
+
6-d vector; an emitted 15-coef bivector generates a per-block rotor R = exp(Ω),
|
| 5 |
+
Ω ∈ so(3,3), and the state evolves by matrix action s_t = R_t · s_{t-1}.
|
| 6 |
+
|
| 7 |
+
Signature η = diag(-1,-1,-1,+1,+1,+1) (matches cl33_t3lm/kge_inference.py).
|
| 8 |
+
so(3,3) = { X : XᵀηX preserved } = { X : ηX antisymmetric }. A basis is
|
| 9 |
+
G_ij = η (e_i e_jᵀ - e_j e_iᵀ) over the 15 pairs i<j — 6 compact rotations
|
| 10 |
+
(both axes same-sign) + 9 boosts (mixed-sign).
|
| 11 |
+
|
| 12 |
+
Reversibility (the interpretability handle): every rotor preserves η, so
|
| 13 |
+
R⁻¹ = η Rᵀ η exactly — no inverse needed. The scan is bit-reversible.
|
| 14 |
+
|
| 15 |
+
Everything here runs in float32 (bf16 diverges — proven on the lattice).
|
| 16 |
+
"""
|
| 17 |
+
from __future__ import annotations
|
| 18 |
+
|
| 19 |
+
import torch
|
| 20 |
+
|
| 21 |
+
ETA = torch.tensor([-1., -1., -1., 1., 1., 1.]) # (6,)
|
| 22 |
+
PAIRS = [(i, j) for i in range(6) for j in range(i + 1, 6)] # 15 pairs
|
| 23 |
+
N_GEN = len(PAIRS) # 15
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def build_generators(device=None, dtype=torch.float32) -> torch.Tensor:
|
| 27 |
+
"""(15, 6, 6) so(3,3) generator matrices G_ij = η(e_i e_jᵀ - e_j e_iᵀ)."""
|
| 28 |
+
eta = ETA.to(device=device, dtype=dtype)
|
| 29 |
+
G = torch.zeros(N_GEN, 6, 6, device=device, dtype=dtype)
|
| 30 |
+
for k, (i, j) in enumerate(PAIRS):
|
| 31 |
+
A = torch.zeros(6, 6, device=device, dtype=dtype)
|
| 32 |
+
A[i, j] = 1.0
|
| 33 |
+
A[j, i] = -1.0
|
| 34 |
+
G[k] = eta[:, None] * A # η @ A (η diagonal)
|
| 35 |
+
return G
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def rotor_from_coefs(coefs: torch.Tensor, G: torch.Tensor) -> torch.Tensor:
|
| 39 |
+
"""coefs (..., 15) -> R (..., 6, 6) = matrix_exp(Σ coefs·G).
|
| 40 |
+
|
| 41 |
+
Runs in the coefs' dtype; caller must keep this in float32.
|
| 42 |
+
"""
|
| 43 |
+
Omega = torch.einsum("...k,kij->...ij", coefs, G) # (...,6,6) in so(3,3)
|
| 44 |
+
return torch.linalg.matrix_exp(Omega)
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
def rotor_inverse(R: torch.Tensor) -> torch.Tensor:
|
| 48 |
+
"""R⁻¹ = η Rᵀ η — exact, no solve. R (...,6,6)."""
|
| 49 |
+
eta = ETA.to(device=R.device, dtype=R.dtype)
|
| 50 |
+
Rt = R.transpose(-1, -2)
|
| 51 |
+
return eta[..., :, None] * Rt * eta[..., None, :]
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
def eta_inner(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
|
| 55 |
+
"""η-metric inner product ⟨a,b⟩_η over the last dim (6). a,b (...,6)."""
|
| 56 |
+
eta = ETA.to(device=a.device, dtype=a.dtype)
|
| 57 |
+
return (a * eta * b).sum(-1)
|
| 58 |
+
|
| 59 |
+
|
| 60 |
+
def q_invariant(v: torch.Tensor) -> torch.Tensor:
|
| 61 |
+
"""Q(v) = -v0²-v1²-v2²+v3²+v4²+v5² (preserved by every rotor)."""
|
| 62 |
+
return eta_inner(v, v)
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
__all__ = ["ETA", "PAIRS", "N_GEN", "build_generators", "rotor_from_coefs",
|
| 66 |
+
"rotor_inverse", "eta_inner", "q_invariant"]
|
t3v3_wedge_memory.py
ADDED
|
@@ -0,0 +1,173 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Wedge associative memory — the geometric-algebra-native fix for cl33-opLM's
|
| 2 |
+
key→value binding failure (MQAR recall ≈ 1/KV). See WEDGE_MEMORY_DESIGN.md.
|
| 3 |
+
|
| 4 |
+
A rotor can only rotate the state; it cannot do an outer-product WRITE, which is how
|
| 5 |
+
associative binding is stored (mLSTM: C += v kᵀ, read C q). The outer product in
|
| 6 |
+
geometric algebra is the wedge, so we add a persistent leaky bivector memory:
|
| 7 |
+
|
| 8 |
+
M_t = γ_t · M_{t-1} + (K_t ∧ V_t) (write; bivector, per block)
|
| 9 |
+
r_t = Q_t ⌋ M_{t-1} (read, before write — causal, no self-match)
|
| 10 |
+
|
| 11 |
+
With the Cl(3,3) η-metric, the wedge/contraction reduce to a clean matrix form:
|
| 12 |
+
K∧V as matrix: W = (V Kᵀ − K Vᵀ) · diag(η) so that
|
| 13 |
+
Q ⌋ (K∧V) = W q = ⟨Q,K⟩_η V − ⟨Q,V⟩_η K
|
| 14 |
+
i.e. when the query matches a stored key, the read returns that key's value.
|
| 15 |
+
This is fast-weight / linear-attention associative memory, native to the algebra —
|
| 16 |
+
the O(T) linear-recurrent sibling of the O(T²) operator-attention (same emitted Q/K).
|
| 17 |
+
"""
|
| 18 |
+
from __future__ import annotations
|
| 19 |
+
|
| 20 |
+
import torch
|
| 21 |
+
import torch.nn as nn
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
class WedgeMemory(nn.Module):
|
| 25 |
+
"""Per-block leaky bivector associative memory. Reuses the emitted query/key
|
| 26 |
+
rotors (via Rq, Rk applied to the state) — no new emission. Returns the read
|
| 27 |
+
r_t (grade-1, 6-d per block) to concatenate into the readout features.
|
| 28 |
+
|
| 29 |
+
Value = the transported state s_t (algebra-pure; keeps the operator-only
|
| 30 |
+
bottleneck / transparency — see design §5/§6 for the richer-V escalation)."""
|
| 31 |
+
|
| 32 |
+
def __init__(self, gate_bias_init: float = 3.0, transport: bool = True,
|
| 33 |
+
key_source: str = "operator", n_gen: int = 15, induction: bool = True,
|
| 34 |
+
d_model: int = None, nb: int = None, value_source: str = "same",
|
| 35 |
+
delta_rule: bool = False):
|
| 36 |
+
super().__init__()
|
| 37 |
+
self.induction = induction # key ← PREVIOUS token/operator (adjacency for MQAR)
|
| 38 |
+
self.key_source = key_source
|
| 39 |
+
# value_source: "same" = value from the same source as the key (default,
|
| 40 |
+
# legacy). "operator" = value is a learned 6-d projection of the emitted
|
| 41 |
+
# state-operator Bs (algebra-only, ablatable) — the G1 experiment: token KEY
|
| 42 |
+
# (clean context-free address) + ALGEBRA VALUE (LINEAR_TRANSPARENT_MEMORY.md).
|
| 43 |
+
self.value_source = value_source
|
| 44 |
+
self.delta_rule = delta_rule # DeltaNet residual write: store v - M·k (reduces
|
| 45 |
+
# write interference by construction) — G/D3.
|
| 46 |
+
if value_source == "operator":
|
| 47 |
+
self.proj_v_op = nn.Linear(n_gen, 6)
|
| 48 |
+
self.nb = nb
|
| 49 |
+
# input-dependent forget gate γ_t = σ(w·x + b). bias_init 3.0 → γ ≈ 0.95.
|
| 50 |
+
# key_source options:
|
| 51 |
+
# "operator": Q/K/V from emitted operators B_q/B_k/B_state (per-token but
|
| 52 |
+
# CONTEXT-MIXED by the causal emitter → induction can't form a
|
| 53 |
+
# clean key; falsified, recall ≈ 1/KV).
|
| 54 |
+
# "state": from the transported state (frame-entangled; + adjoint transport).
|
| 55 |
+
# "token": grade-1 Cl(3,3) projection of the RAW token embedding —
|
| 56 |
+
# CONTEXT-FREE → clean induction key. Fully-transparent
|
| 57 |
+
# (algebra-native η-wedge read) version of the token-copy win.
|
| 58 |
+
if key_source == "operator":
|
| 59 |
+
self.proj_q = nn.Linear(n_gen, 6); self.proj_k = nn.Linear(n_gen, 6)
|
| 60 |
+
self.proj_v = nn.Linear(n_gen, 6)
|
| 61 |
+
gate_in = n_gen
|
| 62 |
+
elif key_source == "token":
|
| 63 |
+
self.proj_q = nn.Linear(d_model, nb * 6); self.proj_k = nn.Linear(d_model, nb * 6)
|
| 64 |
+
self.proj_v = nn.Linear(d_model, nb * 6)
|
| 65 |
+
gate_in = d_model
|
| 66 |
+
else:
|
| 67 |
+
gate_in = 6
|
| 68 |
+
self.gate = nn.Linear(gate_in, 1)
|
| 69 |
+
nn.init.zeros_(self.gate.weight)
|
| 70 |
+
nn.init.constant_(self.gate.bias, gate_bias_init)
|
| 71 |
+
self.transport = transport and key_source == "state"
|
| 72 |
+
|
| 73 |
+
def forward(self, S, Rq, Rk, eta, R_state=None, Bs=None, Bq=None, Bk=None, tok_emb=None):
|
| 74 |
+
"""Returns r (B,T,nb,6). fp32. operator: Bq/Bk/Bs (B,T,nb,15); state: S/Rq/Rk;
|
| 75 |
+
token: tok_emb (B,T,d_model) → grade-1 projections."""
|
| 76 |
+
eta = eta.float()
|
| 77 |
+
B, T, nb = S.shape[:3]
|
| 78 |
+
if self.key_source == "operator":
|
| 79 |
+
Q = self.proj_q(Bq.float()); K = self.proj_k(Bk.float())
|
| 80 |
+
V = self.proj_v(Bs.float()) # per-token, context-mixed
|
| 81 |
+
gamma = torch.sigmoid(self.gate(Bs.float())).squeeze(-1)
|
| 82 |
+
elif self.key_source == "token":
|
| 83 |
+
te = tok_emb.float(); nb = self.nb
|
| 84 |
+
Q = self.proj_q(te).view(B, T, nb, 6) # context-free grade-1 keys
|
| 85 |
+
K = self.proj_k(te).view(B, T, nb, 6)
|
| 86 |
+
if self.value_source == "operator":
|
| 87 |
+
# G1: algebra-only value — a learned projection of the emitted
|
| 88 |
+
# state-operator. Ablating operators (Bs=0) zeroes the value content
|
| 89 |
+
# → read dies → zero bypass (the property token-value lacked).
|
| 90 |
+
V = self.proj_v_op(Bs.float()).view(B, T, nb, 6)
|
| 91 |
+
else:
|
| 92 |
+
V = self.proj_v(te).view(B, T, nb, 6) # token content (bypass)
|
| 93 |
+
gamma = torch.sigmoid(self.gate(te)).expand(B, T, nb)
|
| 94 |
+
else:
|
| 95 |
+
S = S.float(); Rq = Rq.float(); Rk = Rk.float()
|
| 96 |
+
Q = torch.einsum("btnij,btnj->btni", Rq, S)
|
| 97 |
+
K = torch.einsum("btnij,btnj->btni", Rk, S)
|
| 98 |
+
V = S
|
| 99 |
+
gamma = torch.sigmoid(self.gate(S)).squeeze(-1)
|
| 100 |
+
if self.induction:
|
| 101 |
+
# key at t ← the PREVIOUS token/operator → value_t stored under its
|
| 102 |
+
# predecessor, so a query retrieves what FOLLOWED it (MQAR adjacency).
|
| 103 |
+
K = torch.cat([torch.zeros_like(K[:, :1]), K[:, :-1]], dim=1)
|
| 104 |
+
etad = eta[None, None, :] # (1,1,6) for R⁻¹ = ηRᵀη
|
| 105 |
+
if self.transport and R_state is not None:
|
| 106 |
+
R_state = R_state.float()
|
| 107 |
+
|
| 108 |
+
M = torch.zeros(B, nb, 6, 6, device=S.device, dtype=torch.float32)
|
| 109 |
+
reads = []
|
| 110 |
+
for t in range(T):
|
| 111 |
+
if self.transport and R_state is not None:
|
| 112 |
+
Rt = R_state[:, t]
|
| 113 |
+
Rt_inv = etad[..., None] * Rt.transpose(-1, -2) * etad[..., None, :]
|
| 114 |
+
M = Rt @ M @ Rt_inv
|
| 115 |
+
q_t = Q[:, t] # (B,nb,6)
|
| 116 |
+
# READ before write (causal; avoids the query matching its own write)
|
| 117 |
+
reads.append(torch.einsum("bnij,bnj->bni", M, q_t))
|
| 118 |
+
# WRITE: W = (v_eff K_tᵀ − K_t v_effᵀ)·diag(η) = the bivector k∧v_eff as a matrix
|
| 119 |
+
v_t, k_t = V[:, t], K[:, t]
|
| 120 |
+
if self.delta_rule:
|
| 121 |
+
# DeltaNet residual: subtract what M already returns for THIS key, so the
|
| 122 |
+
# write corrects rather than clobbers (interference reduction by construction).
|
| 123 |
+
# STABILIZERS (2026-07-11): L2-normalize key (bounds ||M.k||) + beta write-gate
|
| 124 |
+
# in (0,1) => contraction guaranteed, fixes the M-blowup NaN.
|
| 125 |
+
k_t = k_t / (k_t.norm(dim=-1, keepdim=True) + 1e-6)
|
| 126 |
+
vpred = torch.einsum("bnij,bnj->bni", M, k_t)
|
| 127 |
+
v_eff = 0.5 * (v_t - vpred)
|
| 128 |
+
else:
|
| 129 |
+
v_eff = v_t
|
| 130 |
+
W = (torch.einsum("bni,bnj->bnij", v_eff, k_t)
|
| 131 |
+
- torch.einsum("bni,bnj->bnij", k_t, v_eff)) * eta[None, None, None, :]
|
| 132 |
+
g = gamma[:, t][..., None, None] # (B,nb,1,1)
|
| 133 |
+
M = g * M + W
|
| 134 |
+
return torch.stack(reads, dim=1) # (B,T,nb,6)
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
class TokenCopyMemory(nn.Module):
|
| 138 |
+
"""Minimal token-content value/copy channel — the measured transparency cost of
|
| 139 |
+
KV recall. A per-block fast-weight memory keyed on the RAW token embedding
|
| 140 |
+
(context-free AND position-free → the same token is the same key everywhere, so
|
| 141 |
+
it is matchable across positions and the value TOKEN can be copied). This
|
| 142 |
+
deliberately reads token content that bypasses the operator path; how much the
|
| 143 |
+
model leans on it (vs the operators) is the transparency price, measured directly.
|
| 144 |
+
Plain outer-product fast weights (token embeddings are not in the Cl(3,3) metric)."""
|
| 145 |
+
|
| 146 |
+
def __init__(self, d_model: int, nb: int, key_dim: int = 6, gate_bias_init: float = 3.0,
|
| 147 |
+
induction: bool = True):
|
| 148 |
+
super().__init__()
|
| 149 |
+
self.nb, self.kd = nb, key_dim
|
| 150 |
+
self.induction = induction # key ← PREVIOUS token (the shift that makes MQAR
|
| 151 |
+
# solvable: store value under the key that preceded it)
|
| 152 |
+
self.qkv = nn.Linear(d_model, nb * key_dim * 3)
|
| 153 |
+
self.gate = nn.Linear(d_model, nb)
|
| 154 |
+
nn.init.constant_(self.gate.bias, gate_bias_init)
|
| 155 |
+
|
| 156 |
+
def forward(self, tok_emb):
|
| 157 |
+
"""tok_emb (B,T,d_model) — RAW token embedding (no pos). Returns r (B,T,nb,kd)."""
|
| 158 |
+
tok_emb = tok_emb.float()
|
| 159 |
+
B, T, _ = tok_emb.shape
|
| 160 |
+
qkv = self.qkv(tok_emb).view(B, T, self.nb, 3, self.kd)
|
| 161 |
+
q, k, v = qkv[..., 0, :], qkv[..., 1, :], qkv[..., 2, :] # (B,T,nb,kd)
|
| 162 |
+
if self.induction:
|
| 163 |
+
# key at t ← the PREVIOUS token → value_t is stored under its predecessor,
|
| 164 |
+
# so querying with a token retrieves what FOLLOWED it (induction head).
|
| 165 |
+
k = torch.cat([torch.zeros_like(k[:, :1]), k[:, :-1]], dim=1)
|
| 166 |
+
gamma = torch.sigmoid(self.gate(tok_emb)) # (B,T,nb)
|
| 167 |
+
M = torch.zeros(B, self.nb, self.kd, self.kd, device=tok_emb.device, dtype=torch.float32)
|
| 168 |
+
reads = []
|
| 169 |
+
for t in range(T):
|
| 170 |
+
reads.append(torch.einsum("bnij,bnj->bni", M, q[:, t])) # read before write
|
| 171 |
+
outer = torch.einsum("bni,bnj->bnij", v[:, t], k[:, t]) # v ⊗ k
|
| 172 |
+
M = gamma[:, t][..., None, None] * M + outer
|
| 173 |
+
return torch.stack(reads, dim=1) # (B,T,nb,kd)
|
tape_memory.py
ADDED
|
@@ -0,0 +1,180 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Reversible tape memory — forward-readable exact memory native to the algebra.
|
| 2 |
+
See REVERSIBLE_TAPE_DESIGN.md (2026-07-09).
|
| 3 |
+
|
| 4 |
+
Design principle: TOKEN CONTENT MAY SELECT; ONLY ALGEBRA MAY FLOW.
|
| 5 |
+
- Addressing: context-free token keys (the wedge-validated mechanism) produce a
|
| 6 |
+
causal softmax α over past tape positions. Token content touches ONLY α.
|
| 7 |
+
- Read: r_t = Σ_i α_{t,i} · f(tape_i), where f yields pure algebra objects.
|
| 8 |
+
|
| 9 |
+
The tape is the model's own reversible trace. The exact prefix products
|
| 10 |
+
P_t = R_t···R_1 give the relative transport between any two positions:
|
| 11 |
+
|
| 12 |
+
R_{t←i} = P_t · P_i⁻¹, with P_i⁻¹ = η P_iᵀ η (exact, no replay)
|
| 13 |
+
|
| 14 |
+
Degeneracy (design §3): transported absolute states are trivial (R_{t←i}s_i ∝ s_t),
|
| 15 |
+
so the non-degenerate per-position content is exactly:
|
| 16 |
+
V1 "displacement": R_{t←i} applied to a fixed reference → trajectory geometry only
|
| 17 |
+
V2 "state": s_i raw (frame-mismatched on purpose) → algebra-pure past state
|
| 18 |
+
V3 "increment": B_i, the emitted bivector at i → the local operator event
|
| 19 |
+
(token-decodability ~82% per reversibility.py M2 → carries the
|
| 20 |
+
content MQAR needs, without token embeddings in the value path)
|
| 21 |
+
|
| 22 |
+
Op-ablation behavior (transparency mechanics): zeroed operators ⇒ R=I ⇒ P_t=I,
|
| 23 |
+
displacement=const, increment=0, state=s0 — every value mode collapses while the
|
| 24 |
+
token keys survive. The read CONTENT is operator-derived by construction.
|
| 25 |
+
"""
|
| 26 |
+
from __future__ import annotations
|
| 27 |
+
|
| 28 |
+
import torch
|
| 29 |
+
import torch.nn as nn
|
| 30 |
+
import torch.nn.functional as F
|
| 31 |
+
|
| 32 |
+
VALUE_DIMS = {"displacement": 6, "state": 6, "increment": 15,
|
| 33 |
+
"multi": 21} # displacement(6) + increment(15), one shared address
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class TapeMemory(nn.Module):
|
| 37 |
+
"""Content-addressed read over the reversible tape. O(T²) score like the main
|
| 38 |
+
attention (fine at MQAR/32M scale; anchors+subsampling are the long-context
|
| 39 |
+
lever, not needed here)."""
|
| 40 |
+
|
| 41 |
+
def __init__(self, d_model: int, nb: int, value_mode: str = "increment",
|
| 42 |
+
induction: bool = True, dual_address: bool = False, addr: str = "token"):
|
| 43 |
+
super().__init__()
|
| 44 |
+
assert value_mode in VALUE_DIMS, value_mode
|
| 45 |
+
assert addr in ("token", "operator", "target"), addr
|
| 46 |
+
self.nb, self.value_mode, self.induction, self.addr = nb, value_mode, induction, addr
|
| 47 |
+
# context-free token keys — addressing power proven by the wedge token mode
|
| 48 |
+
self.proj_q = nn.Linear(d_model, nb * 6)
|
| 49 |
+
self.proj_k = nn.Linear(d_model, nb * 6)
|
| 50 |
+
# v2.8 OPERATOR addressing (Garret): the address query is the LM's emitted OPERATOR, not the
|
| 51 |
+
# token. Keys = past operators. "what stored operator-combination fits what I'm computing now"
|
| 52 |
+
# instead of "what token matches" — the native relational address (op-dep memory), NOT the
|
| 53 |
+
# redundant token channel that CE routes around (v2.6 co-option). Ablating operators kills BOTH
|
| 54 |
+
# address and value → the whole tape becomes operator-derived (zero token bypass in addressing).
|
| 55 |
+
if addr == "operator":
|
| 56 |
+
self.proj_q_op = nn.Linear(nb * 15, nb * 6)
|
| 57 |
+
self.proj_k_op = nn.Linear(nb * 15, nb * 6)
|
| 58 |
+
# v2.8 TARGET-DRIVEN navigation (the diamond, facet 2): select by GOAL-vs-EFFECT fit, not
|
| 59 |
+
# key-match. target = learned "what I want" from the current operator; offer = each record's
|
| 60 |
+
# rotor ACTION on a learned reference (its composition effect). score = <target, offer>.
|
| 61 |
+
if addr == "target":
|
| 62 |
+
self.proj_target = nn.Linear(nb * 15, nb * 6) # goal, from the current operator Bq
|
| 63 |
+
self.nav_ref = nn.Parameter(torch.randn(nb, 6) * 0.5) # reference the record's rotor acts on
|
| 64 |
+
# v2.5 DUAL-ADDRESS (TAPE_ADDRESSING_V25.md): a token-FAITHFUL exact-match channel
|
| 65 |
+
# alongside the learned proj address. Probes showed the learned proj gets co-opted
|
| 66 |
+
# for LM context on natural text (attribution 0.98 raw-token → 0.006 learned); a
|
| 67 |
+
# dedicated raw-token-cosine address recovers ~0.98. exact_gate zero-init → a fresh
|
| 68 |
+
# v2.5 model is byte-identical to v2.4 at step 0, and v2.4 ckpts load unaffected.
|
| 69 |
+
self.dual_address = dual_address
|
| 70 |
+
if dual_address:
|
| 71 |
+
self.exact_tau = nn.Parameter(torch.tensor(8.0)) # sharp init (blockade-like)
|
| 72 |
+
self.exact_gate = nn.Parameter(torch.zeros(nb, 1)) # per-block, zero-init
|
| 73 |
+
# V1's fixed reference direction per block (constant → carries no content;
|
| 74 |
+
# the read then reflects trajectory displacement only). Only displacement-
|
| 75 |
+
# bearing modes use it — don't create an unused Parameter otherwise.
|
| 76 |
+
self.ref = (nn.Parameter(torch.randn(nb, 6) * 0.5)
|
| 77 |
+
if value_mode in ("displacement", "multi") else None)
|
| 78 |
+
|
| 79 |
+
def out_dim(self) -> int:
|
| 80 |
+
return VALUE_DIMS[self.value_mode]
|
| 81 |
+
|
| 82 |
+
def forward(self, S, Bs, R_state, tok_emb, eta, Bq=None,
|
| 83 |
+
S_tape=None, Bs_tape=None, R_tape=None, tok_tape=None):
|
| 84 |
+
"""S (B,T,nb,6) states; Bs (B,T,nb,15) emitted bivectors; R_state
|
| 85 |
+
(B,T,nb,6,6) per-step rotors; tok_emb (B,T,d_model) RAW token embedding
|
| 86 |
+
(addressing only). Returns r (B,T,nb,out_dim), fp32.
|
| 87 |
+
|
| 88 |
+
*_tape kwargs (SELF_STRUCTURE_CLAIMS.md, Claim-2 surgery): optional
|
| 89 |
+
COUNTERFACTUAL tape record — the read head sees these instead of the
|
| 90 |
+
factual history, while the live computation (scan/attention, and the
|
| 91 |
+
queries q) stays factual. Keys k and values are drawn from the tape
|
| 92 |
+
record (they are part of the record being counterfactually presented).
|
| 93 |
+
Absent kwargs = factual behavior, byte-identical to before."""
|
| 94 |
+
te = tok_emb.float()
|
| 95 |
+
B, T, nb = S.shape[:3]
|
| 96 |
+
S = S.float()
|
| 97 |
+
eta = eta.float()
|
| 98 |
+
# counterfactual record substitution (read-side only)
|
| 99 |
+
S_rec = S if S_tape is None else S_tape.float()
|
| 100 |
+
Bs_rec = Bs if Bs_tape is None else Bs_tape
|
| 101 |
+
R_rec = R_state if R_tape is None else R_tape
|
| 102 |
+
te_rec = te if tok_tape is None else tok_tape.float()
|
| 103 |
+
|
| 104 |
+
# --- addressing ---
|
| 105 |
+
# queries from the FACTUAL present; keys from the (possibly counterfactual) tape record.
|
| 106 |
+
# OPERATOR mode: address by the emitted operator (native relational key); token content
|
| 107 |
+
# never touches the address. TOKEN mode: the original context-free token key.
|
| 108 |
+
if self.addr == "target":
|
| 109 |
+
# goal-vs-effect fit (navigation, NOT key-match; no induction shift — this is not adjacency).
|
| 110 |
+
assert Bq is not None and R_rec is not None, "target addressing needs Bq and R_state (rotor of Bs)"
|
| 111 |
+
target = self.proj_target(Bq.float().reshape(B, T, nb * 15)).view(B, T, nb, 6) # what I want
|
| 112 |
+
ref = self.nav_ref / self.nav_ref.norm(dim=-1, keepdim=True).clamp_min(1e-6) # (nb,6)
|
| 113 |
+
offer = torch.einsum("btnij,nj->btni", R_rec.float(), ref) # record EFFECT = rotor·ref (B,T,nb,6)
|
| 114 |
+
scores = torch.einsum("btni,bsni->bnts", target, offer) / (6 ** 0.5) # <goal_t, effect_s>
|
| 115 |
+
else:
|
| 116 |
+
if self.addr == "operator":
|
| 117 |
+
assert Bq is not None, "operator addressing needs the query operator Bq"
|
| 118 |
+
q = self.proj_q_op(Bq.float().reshape(B, T, nb * 15)).view(B, T, nb, 6)
|
| 119 |
+
k = self.proj_k_op(Bs_rec.float().reshape(B, T, nb * 15)).view(B, T, nb, 6)
|
| 120 |
+
else:
|
| 121 |
+
q = self.proj_q(te).view(B, T, nb, 6)
|
| 122 |
+
k = self.proj_k(te_rec).view(B, T, nb, 6)
|
| 123 |
+
if self.induction:
|
| 124 |
+
# key at i ← the PREVIOUS token: querying with a key-token selects the position of what
|
| 125 |
+
# FOLLOWED it (MQAR adjacency, as in the wedge). Key-match only.
|
| 126 |
+
k = torch.cat([torch.zeros_like(k[:, :1]), k[:, :-1]], dim=1)
|
| 127 |
+
scores = torch.einsum("btni,bsni->bnts", q, k) / (6 ** 0.5) # (B,nb,T,T)
|
| 128 |
+
causal = torch.triu(torch.ones(T, T, device=S.device, dtype=torch.bool), 0)
|
| 129 |
+
scores = scores.masked_fill(causal, float("-inf")) # STRICT past (i<t)
|
| 130 |
+
A = F.softmax(scores, dim=-1)
|
| 131 |
+
A = torch.nan_to_num(A, nan=0.0) # t=0 has no past
|
| 132 |
+
|
| 133 |
+
# v2.5 DUAL-ADDRESS: token-FAITHFUL exact-match channel (raw token cosine, sharp).
|
| 134 |
+
# It only SELECTS (like A) — the value path is unchanged, so transparency holds.
|
| 135 |
+
A_exact = None
|
| 136 |
+
if self.dual_address:
|
| 137 |
+
tq = F.normalize(te, dim=-1) # (B,T,d) raw query token
|
| 138 |
+
tk = F.normalize(te_rec, dim=-1) # keys from the record
|
| 139 |
+
if self.induction:
|
| 140 |
+
tk = torch.cat([torch.zeros_like(tk[:, :1]), tk[:, :-1]], dim=1)
|
| 141 |
+
esc = torch.einsum("btd,bsd->bts", tq, tk) * self.exact_tau
|
| 142 |
+
esc = esc.masked_fill(causal[None], float("-inf"))
|
| 143 |
+
A_exact = torch.nan_to_num(F.softmax(esc, dim=-1), nan=0.0)[:, None] # (B,1,T,T)
|
| 144 |
+
|
| 145 |
+
def _dual(r_ctx, value):
|
| 146 |
+
if A_exact is None:
|
| 147 |
+
return r_ctx
|
| 148 |
+
r_ex = torch.einsum("bnts,bsni->btni", A_exact.expand(B, nb, T, T), value)
|
| 149 |
+
return r_ctx + self.exact_gate * r_ex # gate zero-init → ≡ v2.4
|
| 150 |
+
|
| 151 |
+
# --- read (only algebra flows) ---
|
| 152 |
+
if self.value_mode == "state":
|
| 153 |
+
return _dual(torch.einsum("bnts,bsni->btni", A, S_rec), S_rec)
|
| 154 |
+
if self.value_mode == "increment":
|
| 155 |
+
return _dual(torch.einsum("bnts,bsni->btni", A, Bs_rec.float()), Bs_rec.float())
|
| 156 |
+
if self.value_mode == "multi":
|
| 157 |
+
# one α, two reads: increment (15) + displacement (6) — the local-order
|
| 158 |
+
# and global-path channels together (v2.4 design)
|
| 159 |
+
r_inc = _dual(torch.einsum("bnts,bsni->btni", A, Bs_rec.float()), Bs_rec.float())
|
| 160 |
+
r_disp = self._displacement_read(A, R_rec.float(), S.device, B, T, nb, eta)
|
| 161 |
+
return torch.cat([r_inc, r_disp], dim=-1)
|
| 162 |
+
|
| 163 |
+
# displacement: r_t = P_t · Σ_i α_{t,i} (P_i⁻¹ · ref̂)
|
| 164 |
+
return self._displacement_read(A, R_rec.float(), S.device, B, T, nb, eta)
|
| 165 |
+
|
| 166 |
+
def _displacement_read(self, A, R, device, B, T, nb, eta):
|
| 167 |
+
P = torch.empty(B, T, nb, 6, 6, device=device, dtype=torch.float32)
|
| 168 |
+
acc = torch.eye(6, device=device).expand(B, nb, 6, 6).contiguous()
|
| 169 |
+
for t in range(T):
|
| 170 |
+
acc = R[:, t] @ acc
|
| 171 |
+
P[:, t] = acc
|
| 172 |
+
ref = self.ref / self.ref.norm(dim=-1, keepdim=True).clamp_min(1e-6) # (nb,6)
|
| 173 |
+
# P_i⁻¹ = η P_iᵀ η
|
| 174 |
+
Pinv = eta[None, None, None, :, None] * P.transpose(-1, -2) * eta[None, None, None, None, :]
|
| 175 |
+
u = torch.einsum("btnij,nj->btni", Pinv, ref) # (B,T,nb,6)
|
| 176 |
+
mix = torch.einsum("bnts,bsni->btni", A, u)
|
| 177 |
+
return torch.einsum("btnij,btnj->btni", P, mix)
|
| 178 |
+
|
| 179 |
+
|
| 180 |
+
__all__ = ["TapeMemory", "VALUE_DIMS"]
|
wedge.py
ADDED
|
@@ -0,0 +1,72 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Exterior-product (wedge) tensors for the grade tower — metric-free.
|
| 2 |
+
|
| 3 |
+
The lattice grade tower is g3=⟨B·v⟩₃, g4=⟨B·B_prev⟩₄, g5=⟨B·B_prev·v⟩₅,
|
| 4 |
+
g6=⟨B·B_prev·B_prev2⟩₆. The top-grade part of a product of blades IS their
|
| 5 |
+
wedge, and the wedge (exterior product) is METRIC-FREE — pure antisymmetrization.
|
| 6 |
+
So the grade tower needs no Cayley/signature convention: fixed wedge tensors
|
| 7 |
+
built from blade combinatorics, computed in PARALLEL per position (not in the
|
| 8 |
+
sequential scan → no speed cost).
|
| 9 |
+
|
| 10 |
+
Blade ordering: grade-k blade = sorted k-subset of {0..5} in `combinations`
|
| 11 |
+
order. Grade-2 order matches so33.PAIRS (both are combinations(range(6),2)), so
|
| 12 |
+
the model's emitted 15 bivector coefs align with these tensors directly.
|
| 13 |
+
"""
|
| 14 |
+
from __future__ import annotations
|
| 15 |
+
|
| 16 |
+
from itertools import combinations
|
| 17 |
+
|
| 18 |
+
import torch
|
| 19 |
+
|
| 20 |
+
GRADE_DIM = {0: 1, 1: 6, 2: 15, 3: 20, 4: 15, 5: 6, 6: 1}
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def blades(grade):
|
| 24 |
+
return list(combinations(range(6), grade))
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def _shuffle_sign(A, B):
|
| 28 |
+
"""Sign of wedge of sorted disjoint tuples A,B = (-1)^#{(a,b): a>b}."""
|
| 29 |
+
if set(A) & set(B):
|
| 30 |
+
return 0, None
|
| 31 |
+
inv = sum(1 for a in A for b in B if a > b)
|
| 32 |
+
merged = tuple(sorted(A + B))
|
| 33 |
+
return (-1) ** inv, merged
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def wedge_tensor(ga: int, gb: int) -> torch.Tensor:
|
| 37 |
+
"""(dim_{ga+gb}, dim_ga, dim_gb): C[c] = Σ W[c,a,b]·A[a]·B[b] realizes A∧B."""
|
| 38 |
+
A, B, C = blades(ga), blades(gb), blades(ga + gb)
|
| 39 |
+
Ci = {b: i for i, b in enumerate(C)}
|
| 40 |
+
W = torch.zeros(len(C), len(A), len(B))
|
| 41 |
+
for ia, a in enumerate(A):
|
| 42 |
+
for ib, b in enumerate(B):
|
| 43 |
+
s, m = _shuffle_sign(a, b)
|
| 44 |
+
if s != 0:
|
| 45 |
+
W[Ci[m], ia, ib] = float(s)
|
| 46 |
+
return W
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
class GradeTower(torch.nn.Module):
|
| 50 |
+
"""Given per-position grade-1 state v and emitted grade-2 operators B (with
|
| 51 |
+
causal shifts B_prev, B_prev2), produce the grade tower features.
|
| 52 |
+
|
| 53 |
+
Inputs (all (..., dim)): v (…,6), B (…,15) [B_prev/B_prev2 are B shifted]
|
| 54 |
+
Output: dict of g2..g6 features. g1=v and g0 handled by caller.
|
| 55 |
+
"""
|
| 56 |
+
|
| 57 |
+
def __init__(self):
|
| 58 |
+
super().__init__()
|
| 59 |
+
self.register_buffer("W_2_1", wedge_tensor(2, 1), persistent=False) # B∧v -> g3
|
| 60 |
+
self.register_buffer("W_2_2", wedge_tensor(2, 2), persistent=False) # B∧B -> g4
|
| 61 |
+
self.register_buffer("W_4_1", wedge_tensor(4, 1), persistent=False) # g4∧v -> g5
|
| 62 |
+
self.register_buffer("W_4_2", wedge_tensor(4, 2), persistent=False) # g4∧B2-> g6
|
| 63 |
+
|
| 64 |
+
def forward(self, v, B, B_prev, B_prev2):
|
| 65 |
+
g3 = torch.einsum("cab,...a,...b->...c", self.W_2_1, B, v) # (...,20)
|
| 66 |
+
g4 = torch.einsum("cab,...a,...b->...c", self.W_2_2, B, B_prev) # (...,15)
|
| 67 |
+
g5 = torch.einsum("cab,...a,...b->...c", self.W_4_1, g4, v) # (...,6)
|
| 68 |
+
g6 = torch.einsum("cab,...a,...b->...c", self.W_4_2, g4, B_prev2) # (...,1)
|
| 69 |
+
return g3, g4, g5, g6
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
__all__ = ["GRADE_DIM", "wedge_tensor", "GradeTower"]
|