AbstractPhil commited on
Commit
3af28f5
Β·
verified Β·
1 Parent(s): e289051

Update experiments/exp_007_aleph_routed_attention/4_aleph_lm.py

Browse files
experiments/exp_007_aleph_routed_attention/4_aleph_lm.py CHANGED
@@ -0,0 +1,1031 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # aleph_lm.py
2
+ """
3
+ AlephLM β€” prediction through the codebook, with guarantees
4
+ ===========================================================
5
+
6
+ The composite reduction (2026-06-09): a causal trigram LM in which the aleph
7
+ signed-projective address is load-bearing at ALL THREE stations β€” input
8
+ addressing, mixing, and prediction. The codebook receives gradient from
9
+ routing, from the predicted next-address pi, and from every candidate address
10
+ kappa. One geometry, closed loop, smooth everywhere (no argmax in the train
11
+ path): differential trigram-to-trigram prediction.
12
+
13
+ CODEC bytes -> trigrams g_t (stride 3)
14
+ EMBED e_t = sum_c E_c[g_t[c]] (byte-factored)
15
+ MIX AlephRoutedAttention hub layers, shared codebook, causal
16
+ PREDICT pi = softmax([w; -w]), w = W_pi h_t (free antipodal-tied)
17
+ CANDIDATE kappa(tau) = address(normalize(W_k sum_c E_c[tau[c]]))
18
+ SCORE logit(tau) = alpha * log( pi+ . kappa+(tau) + pi- . kappa-(tau) )
19
+ OUTPUT hybrid: P(g) = g_in * P_bank(g | in) + (1-g_in) * P_byte(g)
20
+
21
+ THE GUARANTEE LEDGER (all demonstrated numerically 2026-06-09; see session log):
22
+ T1/T2 pi parameterization: address-constrained pi is projectively UNIMODAL
23
+ (logits linear in x-hat) β€” a two-spike target is unreachable (best
24
+ joint mass 1.4% vs 50% needed). The free antipodal-tied simplex
25
+ represents any tied-logit distribution. DEFAULT: free tied simplex;
26
+ address-constrained is the unimodal ablation (pi_mode='address').
27
+ T3 Tied [w; -w] implies p+k * p-k is CONSTANT across k β€” every axis is
28
+ forced to an orientation stance. Feature-or-bug: empirical.
29
+ T4 The 3x256 byte-product head cannot express within-trigram byte
30
+ correlation (rank-1 tensor over 256^3); it is the guaranteed-floor
31
+ baseline (head='byte'), not the main head.
32
+ T5 The hybrid output is a PROPER full-support distribution and its CE
33
+ decomposes exactly: -log P(g) = -log gate_branch - log P_branch(g).
34
+ Implemented verbatim. "Run all three banks" = ablations inside one
35
+ provably-correct machine.
36
+ T6 Raw-score softmax over a bank has a sharpness ceiling (scores in
37
+ (0,1] => CE floor 7.32 nats at M=4096). Logits are LOG-kernel with a
38
+ learnable scale alpha. Non-negotiable.
39
+ T7 Output logit rank <= 2K (softmax bottleneck): K governs attention
40
+ rank, output rank, and mode capacity β€” one knob, three proven roles.
41
+ T8 The write-head target Delta-z = sum of future addresses is the
42
+ ORDER-MARGINALIZED multiset of the next W trigrams (permutation-
43
+ invariant by commutativity). It predicts WHAT comes, not the order.
44
+ Learnability rests on the empirical rank-10 occupancy result.
45
+ Lit. Sampled softmax requires the log-Q correction; with a uniform
46
+ proposal the correction is constant and cancels in the softmax
47
+ (target always included). head='sampled' implements exactly this.
48
+
49
+ THE BRANCHING GAUGE (the [TAU] kernel invariant, inverted): a single unit row
50
+ has conf = ||(p+ - p-)A|| pinned at f(tau,K,D). A PREDICTED pi is not so
51
+ bound β€” implied confidence below the invariant is the model declaring
52
+ superposition. branching_frac is monitored from step zero.
53
+
54
+ Banks: 'corpus' (top-M trigrams of the training stream), 'wordnet'
55
+ (AbstractPhil/wordnet-lexical-topology char_eng_3gram, frequency-ranked,
56
+ filtered to exact 3-byte UTF-8), or per-step 'sampled' negatives.
57
+
58
+ Usage (Blackwell / A100):
59
+ from aleph_lm import AlephLMConfig, train_aleph_lm
60
+ r = train_aleph_lm(AlephLMConfig(steps=10_000, device='cuda',
61
+ head='hybrid', bank_source='wordnet'))
62
+
63
+ Depends: aleph_routed_attention.py, aleph_trigram_lm.py in the same directory.
64
+ Author: AbstractPhil + Mirel Date: 2026-06-09 License: MIT
65
+ """
66
+
67
+ from __future__ import annotations
68
+
69
+ import math
70
+ import os
71
+ import time
72
+ from dataclasses import dataclass
73
+ from typing import Dict, List, Optional, Tuple
74
+
75
+ import numpy as np
76
+ import torch
77
+ import torch.nn as nn
78
+ import torch.nn.functional as F
79
+ from torch import Tensor
80
+
81
+ from aleph_routed_attention import AlephRoutedAttention, AlephAttentionConfig
82
+ from aleph_trigram_lm import TrigramStream, statute
83
+
84
+
85
+ # ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
86
+ # Config
87
+ # ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
88
+
89
+ @dataclass
90
+ class AlephLMConfig:
91
+ # substrate
92
+ corpus_id: str = "wikitext-103-raw-v1"
93
+ split: str = "train"
94
+ max_corpus_bytes: Optional[int] = 100_000_000
95
+ seq_len: int = 256 # trigrams (3*seq_len bytes)
96
+ seed: int = 1234
97
+
98
+ # tower
99
+ dim: int = 384
100
+ n_layers: int = 4
101
+ n_heads: int = 6
102
+ K: int = 64
103
+ D_addr: int = 4
104
+ tau: float = 0.1
105
+ codebook_init: object = "random"
106
+ shared_codebook: bool = True
107
+
108
+ # prediction head (the ledger's resolutions)
109
+ head: str = "hybrid" # 'hybrid' | 'byte' | 'bank' | 'sampled'
110
+ pi_mode: str = "free" # 'free' (T2 default) | 'address' (unimodal ablation)
111
+ bank_source: str = "corpus" # 'corpus' | 'wordnet'
112
+ bank_size: int = 4096
113
+ n_negatives: int = 1024 # head='sampled'
114
+ logit_scale_init: float = 1.0 # alpha on the log-kernel logits (T6)
115
+
116
+ # hybrid bank scorer: 'kernel' (log-kernel, T6) or 'pmix' β€” a mixture of
117
+ # J pointers on S^(d_point-1): logits(c) = logsumexp_j [log w_j + T yhat_j.c]
118
+ # = Mixture-of-Softmaxes in sphere coordinates. Theorem-backed twice over:
119
+ # raises output rank past the T7 bottleneck (MoS, Yang et al.), and gives
120
+ # the pointer J modes so the barycenter pathology (unimodal aim at a
121
+ # multimodal future) is structurally removed. Full-bank softmax retained:
122
+ # T5 propriety intact. PREREGISTERED statute prediction: pmix candidate
123
+ # coords bypass the codebook (W_cand48), removing prediction-side
124
+ # discrimination pressure -> expect dev near the zero group, vs kernel's
125
+ # +0.013. The dose-response gets a within-architecture test.
126
+ bank_scorer: str = "kernel" # 'kernel' | 'pmix'
127
+ # ── Tier-A scaling switches (2026-06-11) ──
128
+ # bank_softmax='sampled' (pmix only): train CE over {batch targets} βˆͺ
129
+ # n_bank_samples uniform negatives with the log-Q correction (Jean et al.;
130
+ # shared-negative approximation documented at the call site). The REPORTED
131
+ # bpb stays honest: at every log step the in-bank NLL is recomputed against
132
+ # the FULL bank in chunked no_grad. pos_mode='clamp' saturates position
133
+ # indices at seq_len-1 -> unbounded streaming/generation (decay-gated state
134
+ # is the principled v2). train_mode='stream' = TBPTT-1 over `segments`
135
+ # carried streaming states: effective context segments*seq_len at constant
136
+ # memory (requires pos_mode='clamp').
137
+ bank_softmax: str = "full" # 'full' | 'sampled'
138
+ n_bank_samples: int = 8192
139
+ pos_mode: str = "absolute" # 'absolute' | 'clamp'
140
+ train_mode: str = "window" # 'window' | 'stream'
141
+ segments: int = 4
142
+ amp: bool = True # bf16 autocast on CUDA
143
+ compile_backbone: bool = False # compile backbone only (tensor out)
144
+ n_pointers: int = 4 # J mixture components (pmix)
145
+
146
+ # pointer head (head='pointer'): NN-on-the-sphere decode
147
+ d_point: int = 48 # pointer sphere dim (band-valid; the
148
+ # capacity table gives the decode
149
+ # budget theta_NN/2 at this D)
150
+ pointer_k: int = 32 # hard negatives = target's k sphere-NN
151
+ pointer_cos_weight: float = 0.5 # aiming regularizer (contrastive CE
152
+ # is the main learner β€” lit. caveat)
153
+ pointer_refresh: int = 200 # steps between NN-table refreshes
154
+ # (candidate coords drift)
155
+
156
+ # write-head (T8, auxiliary multiset prediction)
157
+ write_weight: float = 0.1 # 0 disables
158
+ write_horizon: int = 8 # W: the granularity dial
159
+
160
+ # training
161
+ steps: int = 10_000
162
+ batch_size: int = 32
163
+ accum_steps: int = 1
164
+ lr: float = 3e-4
165
+ lr_decay: bool = True
166
+ div_weight: float = 0.0
167
+ log_every: int = 250
168
+ device: str = "cuda" if torch.cuda.is_available() else "cpu"
169
+
170
+ # outputs
171
+ snapshot_codebook: bool = True
172
+ snapshot_path: str = "aleph_lm_snaps.pt"
173
+ checkpoint_path: Optional[str] = "aleph_lm.pt"
174
+
175
+ def __post_init__(self):
176
+ assert self.head in ("hybrid", "byte", "bank", "sampled", "pointer")
177
+ assert self.pi_mode in ("free", "address")
178
+ assert self.bank_scorer in ("kernel", "pmix")
179
+ assert self.bank_softmax in ("full", "sampled")
180
+ assert self.pos_mode in ("absolute", "clamp")
181
+ assert self.train_mode in ("window", "stream")
182
+ if self.bank_softmax == "sampled":
183
+ assert self.head == "hybrid" and self.bank_scorer == "pmix", \
184
+ "sampled softmax is wired for the pmix hybrid"
185
+ if self.train_mode == "stream":
186
+ assert self.pos_mode == "clamp", "stream training needs pos_mode='clamp'"
187
+ assert self.bank_source in ("corpus", "wordnet") \
188
+ or os.path.isfile(str(self.bank_source)), \
189
+ f"bank_source must be 'corpus'|'wordnet'|path to bank .pt"
190
+ assert self.dim % self.n_heads == 0
191
+ assert self.write_horizon >= 1
192
+ tag = self.head
193
+ if self.head in ("hybrid", "bank"):
194
+ b = (os.path.splitext(os.path.basename(str(self.bank_source)))[0]
195
+ if os.path.isfile(str(self.bank_source)) else self.bank_source)
196
+ tag += f"_{b}"
197
+ if self.pi_mode != "free":
198
+ tag += f"_{self.pi_mode}"
199
+ if self.head == "hybrid" and self.bank_scorer == "pmix":
200
+ tag += f"_pmix{self.n_pointers}"
201
+ if self.checkpoint_path == "aleph_lm.pt":
202
+ self.checkpoint_path = f"aleph_lm_{tag}.pt"
203
+ if self.snapshot_path == "aleph_lm_snaps.pt":
204
+ self.snapshot_path = f"aleph_lm_snaps_{tag}.pt"
205
+
206
+
207
+ # ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
208
+ # Candidate banks
209
+ # ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
210
+
211
+ def _tri_ids(tri: Tensor) -> Tensor:
212
+ """(..., 3) bytes -> scalar trigram id in [0, 256^3)."""
213
+ return tri[..., 0] * 65536 + tri[..., 1] * 256 + tri[..., 2]
214
+
215
+
216
+ def build_corpus_bank(stream: TrigramStream, M: int,
217
+ sample_bytes: int = 6_000_000) -> Tensor:
218
+ """Top-M most frequent trigrams of the training stream. (M, 3) long."""
219
+ n = min(sample_bytes, (len(stream.stream) // 3) * 3)
220
+ tri = stream.stream[:n].reshape(-1, 3)
221
+ ids = (tri[:, 0].astype(np.int64) * 65536 + tri[:, 1].astype(np.int64) * 256
222
+ + tri[:, 2].astype(np.int64))
223
+ uniq, counts = np.unique(ids, return_counts=True)
224
+ top = uniq[np.argsort(counts)[::-1][:M]]
225
+ out = np.stack([top // 65536, (top // 256) % 256, top % 256], axis=-1)
226
+ return torch.from_numpy(out.astype(np.int64))
227
+
228
+
229
+ def build_wordnet_bank(M: int) -> Tensor:
230
+ """char_eng_3gram from AbstractPhil/wordnet-lexical-topology, frequency-
231
+ ranked, filtered to exact 3-byte UTF-8. (M', 3) long, M' <= M."""
232
+ from huggingface_hub import hf_hub_download
233
+ import pyarrow.parquet as pq
234
+ p = hf_hub_download("AbstractPhil/wordnet-lexical-topology",
235
+ "data/char_eng_3gram-00000-of-00001.parquet",
236
+ repo_type="dataset")
237
+ t = pq.read_table(p, columns=["ngram", "rank"]).to_pandas()
238
+ t = t.sort_values("rank")
239
+ rows = []
240
+ for s in t["ngram"]:
241
+ b = str(s).encode("utf-8", errors="ignore")
242
+ if len(b) == 3:
243
+ rows.append([b[0], b[1], b[2]])
244
+ if len(rows) >= M:
245
+ break
246
+ assert rows, "wordnet bank empty after 3-byte filter"
247
+ return torch.tensor(rows, dtype=torch.long)
248
+
249
+
250
+ # ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
251
+ # Model
252
+ # ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
253
+
254
+ class AlephLM(nn.Module):
255
+ """The composite reduction. forward_loss(ids, targets) -> (loss, logs)."""
256
+
257
+ def __init__(self, cfg: AlephLMConfig, bank: Optional[Tensor] = None):
258
+ super().__init__()
259
+ self.cfg = cfg
260
+ d = cfg.dim
261
+
262
+ # ── EMBED (byte-factored; shared with candidate composition) ──
263
+ self.byte_emb = nn.ModuleList([nn.Embedding(256, d) for _ in range(3)])
264
+ self.pos = nn.Parameter(0.02 * torch.randn(1, cfg.seq_len, d))
265
+
266
+ # ── MIX (hub tower, shared codebook) ──
267
+ def make_attn():
268
+ return AlephRoutedAttention(AlephAttentionConfig(
269
+ dim=d, num_heads=cfg.n_heads, mode="hub", K=cfg.K,
270
+ D_addr=cfg.D_addr, tau=cfg.tau, causal=True,
271
+ codebook_init=cfg.codebook_init))
272
+ self.layers = nn.ModuleList([
273
+ nn.ModuleDict({"norm1": nn.LayerNorm(d), "attn": make_attn(),
274
+ "norm2": nn.LayerNorm(d),
275
+ "mlp": nn.Sequential(nn.Linear(d, 4 * d), nn.GELU(),
276
+ nn.Linear(4 * d, d))})
277
+ for _ in range(cfg.n_layers)])
278
+ if cfg.shared_codebook:
279
+ shared = self.layers[0]["attn"].codebook
280
+ for L in self.layers[1:]:
281
+ L["attn"].codebook = shared
282
+ self.norm_f = nn.LayerNorm(d)
283
+
284
+ # ── PREDICT: pi over 2K oriented axes ──
285
+ self.W_pi = nn.Linear(d, cfg.K, bias=True) # tied logits [w; -w]
286
+ if cfg.pi_mode == "address": # unimodal ablation (T2)
287
+ self.W_pi_row = nn.Linear(d, cfg.D_addr, bias=False)
288
+ nn.init.orthogonal_(self.W_pi_row.weight)
289
+
290
+ # ── CANDIDATE: compositional addresses (banked heads) ──
291
+ self.W_kappa = nn.Linear(d, cfg.D_addr, bias=False)
292
+ nn.init.orthogonal_(self.W_kappa.weight)
293
+ self.logit_scale = nn.Parameter(
294
+ torch.tensor(float(cfg.logit_scale_init))) # alpha (T6)
295
+
296
+ # ── byte-product head (T4 floor + hybrid tail) ──
297
+ self.byte_heads = nn.ModuleList([nn.Linear(d, 256) for _ in range(3)])
298
+
299
+ # ── hybrid gate (T5) ──
300
+ self.gate = nn.Linear(d, 1)
301
+
302
+ # ── pmix bank scorer: J-pointer mixture (MoS on the sphere) ──
303
+ if cfg.head == "hybrid" and cfg.bank_scorer == "pmix":
304
+ J = cfg.n_pointers
305
+ self.W_pmix = nn.Linear(d, J * cfg.d_point, bias=False)
306
+ nn.init.orthogonal_(self.W_pmix.weight)
307
+ self.W_mixgate = nn.Linear(d, J)
308
+ self.W_cand48 = nn.Linear(d, cfg.d_point, bias=False)
309
+ nn.init.orthogonal_(self.W_cand48.weight)
310
+ self.point_T = nn.Parameter(torch.tensor(10.0))
311
+
312
+ # ── pointer head: predict a point on S^(d_point-1), decode by NN ──
313
+ if cfg.head == "pointer":
314
+ self.W_point = nn.Linear(d, cfg.d_point, bias=False)
315
+ nn.init.orthogonal_(self.W_point.weight)
316
+ self.W_cand48 = nn.Linear(d, cfg.d_point, bias=False)
317
+ nn.init.orthogonal_(self.W_cand48.weight)
318
+ self.point_T = nn.Parameter(torch.tensor(10.0)) # contrastive inv-temp
319
+ self.register_buffer("_nn_table", torch.zeros(0, dtype=torch.long),
320
+ persistent=False)
321
+ self._nn_step = -1
322
+
323
+ # ── write-head (T8): predicted Delta-z over 2K ──
324
+ if cfg.write_weight > 0:
325
+ self.W_write = nn.Linear(d, 2 * cfg.K)
326
+
327
+ # ── bank registration ──
328
+ if bank is not None:
329
+ self.register_buffer("bank", bank) # (M, 3)
330
+ self.register_buffer("bank_ids_sorted",
331
+ _tri_ids(bank).sort().values) # membership
332
+ self.register_buffer("bank_perm",
333
+ _tri_ids(bank).argsort()) # sorted->orig
334
+ else:
335
+ self.bank = None
336
+
337
+ # ---- shared codebook handle ----
338
+ @property
339
+ def codebook(self) -> Tensor:
340
+ return self.layers[0]["attn"].codebook
341
+
342
+ def aleph_layers(self) -> List[AlephRoutedAttention]:
343
+ return [m for m in self.modules() if isinstance(m, AlephRoutedAttention)]
344
+
345
+ # ---- address of arbitrary unit rows vs the SHARED codebook ----
346
+ def _address_rows(self, rows: Tensor) -> Tuple[Tensor, Tensor]:
347
+ A = F.normalize(self.codebook, dim=-1)
348
+ u = (rows @ A.t()) / self.cfg.tau
349
+ m = u.abs().amax(-1, keepdim=True)
350
+ ep, en = torch.exp(u - m), torch.exp(-u - m)
351
+ Z = (ep + en).sum(-1, keepdim=True)
352
+ return ep / Z, en / Z
353
+
354
+ # ---- PREDICT ----
355
+ def _pi(self, h: Tensor) -> Tuple[Tensor, Tensor]:
356
+ """pi over 2K oriented axes. 'free': tied simplex (T2 default).
357
+ 'address': unimodal ablation."""
358
+ if self.cfg.pi_mode == "address":
359
+ row = F.normalize(self.W_pi_row(h), dim=-1)
360
+ return self._address_rows(row)
361
+ w = self.W_pi(h) # (..., K)
362
+ m = w.abs().amax(-1, keepdim=True)
363
+ ep, en = torch.exp(w - m), torch.exp(-w - m)
364
+ Z = (ep + en).sum(-1, keepdim=True)
365
+ return ep / Z, en / Z
366
+
367
+ # ---- CANDIDATE addresses for a (M, 3) byte bank ----
368
+ def _kappa(self, bank: Tensor) -> Tuple[Tensor, Tensor]:
369
+ e = sum(emb(bank[:, i]) for i, emb in enumerate(self.byte_emb))
370
+ rows = F.normalize(self.W_kappa(e), dim=-1) # (M, D_addr)
371
+ return self._address_rows(rows)
372
+
373
+ # ---- SCORE: log-kernel logits (T6) ----
374
+ def _bank_logits(self, pi_p: Tensor, pi_m: Tensor,
375
+ k_p: Tensor, k_m: Tensor) -> Tensor:
376
+ s = pi_p @ k_p.t() + pi_m @ k_m.t() # (..., M), > 0
377
+ return self.logit_scale * torch.log(s.clamp_min(1e-9))
378
+
379
+ # ---- pmix bank logits: logsumexp over J pointers (MoS on the sphere) ----
380
+ def _pmix_logits(self, h: Tensor,
381
+ cols: Optional[Tensor] = None) -> Tuple[Tensor, Dict[str, float]]:
382
+ cfg = self.cfg
383
+ J = cfg.n_pointers
384
+ coords = self._cand_coords() # (M, d_point)
385
+ if cols is not None:
386
+ coords = coords[cols] # (Mc, d_point)
387
+ y = self.W_pmix(h).view(*h.shape[:-1], J, cfg.d_point)
388
+ y = F.normalize(y, dim=-1) # (B,S,J,dp)
389
+ mix = F.log_softmax(self.W_mixgate(h), dim=-1) # (B,S,J)
390
+ sims = torch.einsum("bsjd,md->bsjm", y, coords) * self.point_T
391
+ logits = torch.logsumexp(mix.unsqueeze(-1) + sims, dim=2) # (B,S,M)
392
+ with torch.no_grad(): # mode diagnostics
393
+ pw = torch.einsum("bsjd,bskd->bsjk", y, y)
394
+ off = pw.masked_select(~torch.eye(J, dtype=torch.bool,
395
+ device=h.device)
396
+ .expand_as(pw)).clamp(-1, 1)
397
+ spread = torch.acos(off).mean().item() * 180 / math.pi
398
+ usage = mix.exp().mean(dim=(0, 1))
399
+ ent = -(usage * usage.clamp_min(1e-9).log()).sum().item() / math.log(J)
400
+ return logits, {"mode_spread_deg": spread, "mix_entropy": ent}
401
+
402
+ # ---- position (absolute, or clamped for unbounded streaming) ----
403
+ def _pos_slice(self, S: int, offset: int = 0) -> Tensor:
404
+ if self.cfg.pos_mode == "clamp":
405
+ idx = (torch.arange(S, device=self.pos.device) + offset
406
+ ).clamp_max(self.cfg.seq_len - 1)
407
+ return self.pos[:, idx]
408
+ return self.pos[:, offset: offset + S]
409
+
410
+ # ---- honest full-bank in-bank NLL (chunked, no_grad) for sampled mode ----
411
+ @torch.no_grad()
412
+ def _pmix_full_nll(self, h: Tensor, bidx: Tensor, in_bank: Tensor) -> Tensor:
413
+ M = self.bank.shape[0]
414
+ lse = None
415
+ tgt_logit = torch.zeros_like(bidx, dtype=h.dtype)
416
+ for lo in range(0, M, 8192):
417
+ cols = torch.arange(lo, min(lo + 8192, M), device=h.device)
418
+ lg, _ = self._pmix_logits(h, cols=cols) # (B,S,Mc)
419
+ chunk_lse = torch.logsumexp(lg, dim=-1)
420
+ lse = chunk_lse if lse is None else torch.logaddexp(lse, chunk_lse)
421
+ hit = in_bank & (bidx >= lo) & (bidx < lo + cols.numel())
422
+ if hit.any():
423
+ tgt_logit[hit] = lg[hit].gather(
424
+ -1, (bidx[hit] - lo).unsqueeze(-1)).squeeze(-1)
425
+ return lse - tgt_logit # NLL (B,S)
426
+
427
+ # ---- tower ----
428
+ def backbone(self, ids: Tensor) -> Tensor:
429
+ x = sum(emb(ids[..., i]) for i, emb in enumerate(self.byte_emb))
430
+ x = x + self._pos_slice(ids.shape[1])
431
+ for L in self.layers:
432
+ x = x + L["attn"](L["norm1"](x))
433
+ x = x + L["mlp"](L["norm2"](x))
434
+ return self.norm_f(x)
435
+
436
+ # ---- byte-product log-probs of given targets (T4 tail) ----
437
+ def _byte_logprob(self, h: Tensor, targets: Tensor) -> Tensor:
438
+ lp = 0.0
439
+ for c, head in enumerate(self.byte_heads):
440
+ lp = lp + F.log_softmax(head(h), dim=-1).gather(
441
+ -1, targets[..., c:c + 1]).squeeze(-1)
442
+ return lp # (B, S)
443
+
444
+ # ---- bank membership: target -> bank index or -1 ----
445
+ def _bank_index(self, targets: Tensor) -> Tensor:
446
+ tid = _tri_ids(targets)
447
+ pos = torch.searchsorted(self.bank_ids_sorted, tid)
448
+ pos = pos.clamp_max(len(self.bank_ids_sorted) - 1)
449
+ hit = self.bank_ids_sorted[pos] == tid
450
+ idx = self.bank_perm[pos]
451
+ return torch.where(hit, idx, torch.full_like(idx, -1))
452
+
453
+ # ---- write-head target (T8): order-marginalized future address mass ----
454
+ @torch.no_grad()
455
+ def _write_target(self, ids: Tensor) -> Tensor:
456
+ """Delta-z over 2K for horizon W at each position (normalized)."""
457
+ cfg = self.cfg
458
+ a0 = self.layers[0]["attn"]
459
+ x = sum(emb(ids[..., i]) for i, emb in enumerate(self.byte_emb))
460
+ x = x + self.pos[:, : ids.shape[1]]
461
+ kh = a0._split_addr(a0.k_addr(self.layers[0]["norm1"](x)),
462
+ ids.shape[0], ids.shape[1])
463
+ pk_p, pk_m = a0._address(kh) # (B,H,S,K)
464
+ p = torch.cat([pk_p, pk_m], dim=-1).mean(dim=1) # (B,S,2K)
465
+ cs = torch.cat([torch.zeros_like(p[:, :1]), p.cumsum(dim=1)], dim=1)
466
+ W = cfg.write_horizon
467
+ B, S, _ = p.shape
468
+ end = torch.arange(S, device=p.device).clamp_max(S - 1)
469
+ lo = cs[:, 1:] # prefix up to t (incl)
470
+ hi = cs[:, torch.clamp(torch.arange(S, device=p.device) + W, max=S)]
471
+ dz = (hi - lo).clamp_min(0)
472
+ valid = (torch.arange(S, device=p.device) + 1 < S) # at least 1 future tok
473
+ dz = dz / dz.sum(-1, keepdim=True).clamp_min(1e-9)
474
+ return dz, valid
475
+
476
+ # ---- pointer head: compositional D=48 candidate coordinates ----
477
+ def _cand_coords(self) -> Tensor:
478
+ e = sum(emb(self.bank[:, i]) for i, emb in enumerate(self.byte_emb))
479
+ return F.normalize(self.W_cand48(e), dim=-1) # (M, d_point)
480
+
481
+ @torch.no_grad()
482
+ def _refresh_nn(self, coords: Tensor, step: int) -> None:
483
+ """Hard-negative table: each candidate's k nearest sphere neighbors
484
+ (excluding self). Refreshed periodically β€” coordinates drift."""
485
+ cos = coords @ coords.t()
486
+ cos.fill_diagonal_(-2.0)
487
+ self._nn_table = cos.topk(self.cfg.pointer_k, dim=-1).indices # (M, k)
488
+ self._nn_step = step
489
+ # decode budget: theta_NN/2 of the CURRENT candidate constellation
490
+ nn_deg = torch.acos(cos.max(dim=-1).values.clamp(-1, 1)) * 180 / math.pi
491
+ self._decode_budget_deg = (nn_deg.median() / 2).item()
492
+
493
+ def _pointer_loss(self, h: Tensor, targets: Tensor,
494
+ step: int) -> Tuple[Tensor, Dict]:
495
+ """NN-on-the-sphere head (T5-chained with the byte tail):
496
+ in-bank: -log gate - log softmax_{target βˆͺ kNN(target)}(T * yhatΒ·c)
497
+ + lambda_cos (1 - yhatΒ·c_target) [aiming term]
498
+ out-bank: -log(1-gate) - log P_byte(g)
499
+ Decode metric: exact-NN rate + median angular error vs the budget
500
+ theta_NN/2 (the decode-correctness theorem)."""
501
+ cfg = self.cfg
502
+ logs: Dict[str, float] = {}
503
+ coords = self._cand_coords() # (M, d_point)
504
+ if step - self._nn_step >= cfg.pointer_refresh or len(self._nn_table) == 0:
505
+ self._refresh_nn(coords.detach(), step)
506
+
507
+ yhat = F.normalize(self.W_point(h), dim=-1) # (B,S,d_point)
508
+ bidx = self._bank_index(targets)
509
+ in_bank = bidx >= 0
510
+ logs["coverage"] = in_bank.float().mean().item()
511
+
512
+ g_logit = self.gate(h).squeeze(-1)
513
+ nll_byte = -self._byte_logprob(h, targets)
514
+
515
+ B, S = bidx.shape
516
+ tgt = bidx.clamp_min(0) # (B,S)
517
+ negs = self._nn_table[tgt] # (B,S,k) hard negatives
518
+ cand_idx = torch.cat([tgt.unsqueeze(-1), negs], dim=-1) # (B,S,1+k)
519
+ c = coords[cand_idx] # (B,S,1+k,d_point)
520
+ logits = torch.einsum("bsd,bsnd->bsn", yhat, c) * self.point_T
521
+ nll_point = F.cross_entropy(
522
+ logits.reshape(-1, logits.shape[-1]),
523
+ torch.zeros(B * S, dtype=torch.long, device=h.device),
524
+ reduction="none").view(B, S)
525
+ cos_t = torch.einsum("bsd,bsd->bs", yhat, coords[tgt])
526
+ aim = cfg.pointer_cos_weight * (1.0 - cos_t)
527
+
528
+ nll = torch.where(in_bank,
529
+ -F.logsigmoid(g_logit) + nll_point + aim,
530
+ -F.logsigmoid(-g_logit) + nll_byte)
531
+ loss = nll.mean()
532
+ logs["bpb"] = loss.item() / 3 / math.log(2)
533
+ logs["gate_acc"] = ((torch.sigmoid(g_logit) > 0.5) == in_bank
534
+ ).float().mean().item()
535
+ with torch.no_grad(): # decode metrics
536
+ if in_bank.any():
537
+ full = (yhat @ coords.t()) # (B,S,M)
538
+ pred = full.argmax(-1)
539
+ logs["nn_exact"] = (pred[in_bank] == tgt[in_bank]
540
+ ).float().mean().item()
541
+ # PROPER eval likelihood: full-bank softmax (comparable to
542
+ # hybrid bpb; the training loss above is contrastive-over-33
543
+ # and is NOT a likelihood β€” do not compare it across heads)
544
+ nll_full = F.cross_entropy(
545
+ (full * self.point_T).reshape(-1, full.shape[-1]),
546
+ tgt.reshape(-1), reduction="none").view_as(tgt)
547
+ nll_eval = torch.where(in_bank,
548
+ -F.logsigmoid(g_logit) + nll_full,
549
+ -F.logsigmoid(-g_logit) + nll_byte)
550
+ logs["bpb_eval"] = nll_eval.mean().item() / 3 / math.log(2)
551
+ ang = torch.acos(cos_t[in_bank].clamp(-1, 1)) * 180 / math.pi
552
+ logs["ang_err_deg"] = ang.median().item()
553
+ logs["budget_deg"] = self._decode_budget_deg
554
+ logs["in_budget"] = (ang < self._decode_budget_deg
555
+ ).float().mean().item()
556
+ return loss, logs
557
+
558
+ # ---- the loss (T5-exact hybrid + auxiliaries) ----
559
+ def forward_loss(self, ids: Tensor, targets: Tensor,
560
+ step: int = 0, h: Optional[Tensor] = None
561
+ ) -> Tuple[Tensor, Dict]:
562
+ cfg = self.cfg
563
+ if h is None:
564
+ h = self.backbone(ids) # (B,S,d)
565
+ logs: Dict[str, float] = {}
566
+
567
+ if cfg.head == "pointer":
568
+ assert self.bank is not None, "pointer head requires a bank"
569
+ return self._pointer_loss(h, targets, step)
570
+
571
+ if cfg.head == "byte":
572
+ nll = -self._byte_logprob(h, targets) # (B,S)
573
+ loss = nll.mean()
574
+ logs["bpb"] = loss.item() / 3 / math.log(2)
575
+ return loss, logs
576
+
577
+ pi_p, pi_m = self._pi(h) # (B,S,K) each
578
+
579
+ if cfg.head == "sampled":
580
+ # uniform negatives + target; uniform proposal => logQ constant,
581
+ # cancels in softmax (literature requirement satisfied trivially)
582
+ B, S, _ = h.shape
583
+ neg = torch.randint(0, 256, (cfg.n_negatives, 3), device=h.device)
584
+ cand = torch.cat([targets.reshape(-1, 3), neg], dim=0)
585
+ cand_ids, inv = torch.unique(_tri_ids(cand), return_inverse=True)
586
+ uniq = torch.stack([cand_ids // 65536, (cand_ids // 256) % 256,
587
+ cand_ids % 256], dim=-1)
588
+ k_p, k_m = self._kappa(uniq)
589
+ logits = self._bank_logits(pi_p.reshape(-1, cfg.K),
590
+ pi_m.reshape(-1, cfg.K), k_p, k_m)
591
+ tgt_idx = inv[: B * S]
592
+ loss = F.cross_entropy(logits, tgt_idx)
593
+ logs["bpb"] = loss.item() / 3 / math.log(2)
594
+ logs["n_cand"] = float(len(uniq))
595
+ return loss, logs
596
+
597
+ # banked heads
598
+ assert self.bank is not None, "head='hybrid'/'bank' requires a bank"
599
+ bidx = self._bank_index(targets) # (B,S), -1 = miss
600
+ in_bank = bidx >= 0
601
+ logs["coverage"] = in_bank.float().mean().item()
602
+ sampled = (cfg.bank_softmax == "sampled" and self.training)
603
+ if cfg.head == "hybrid" and cfg.bank_scorer == "pmix":
604
+ if sampled:
605
+ # shared-negative sampled softmax: cols = batch targets βˆͺ
606
+ # uniform negatives; log-Q correction log(n/M) on pure
607
+ # negatives only (cols that are some example's target carry
608
+ # Qβ‰ˆ1; tiny documented bias β€” reported bpb is recomputed
609
+ # full-bank below at log steps, so the LOGGED number is exact)
610
+ M = self.bank.shape[0]
611
+ n = min(cfg.n_bank_samples, M)
612
+ tcols = bidx[in_bank].unique()
613
+ samp = torch.randint(0, M, (n,), device=h.device)
614
+ cols = torch.unique(torch.cat([tcols, samp]))
615
+ logits, pm_logs = self._pmix_logits(h, cols=cols)
616
+ logits = logits + torch.where(
617
+ torch.isin(cols, tcols), 0.0,
618
+ -math.log(n / M)).to(logits.dtype) # -(-logQ)
619
+ bidx = torch.searchsorted(cols, bidx.clamp_min(0))
620
+ else:
621
+ logits, pm_logs = self._pmix_logits(h) # (B,S,M)
622
+ logs.update(pm_logs)
623
+ else:
624
+ k_p, k_m = self._kappa(self.bank) # (M,K) each
625
+ logits = self._bank_logits(pi_p, pi_m, k_p, k_m) # (B,S,M)
626
+
627
+ if cfg.head == "bank":
628
+ # ablation head: proper only on covered targets (coverage logged)
629
+ lb = F.log_softmax(logits, dim=-1)
630
+ nll = -lb.gather(-1, bidx.clamp_min(0).unsqueeze(-1)).squeeze(-1)
631
+ loss = nll[in_bank].mean() if in_bank.any() else logits.sum() * 0
632
+ logs["bpb_inbank"] = (loss.item() / 3 / math.log(2)
633
+ if in_bank.any() else float("nan"))
634
+ return loss, logs
635
+
636
+ # ── T5-exact hybrid: -log P(g) per position ──
637
+ g_logit = self.gate(h).squeeze(-1) # (B,S)
638
+ log_g = F.logsigmoid(g_logit)
639
+ log_1mg = F.logsigmoid(-g_logit)
640
+ lb = F.log_softmax(logits, dim=-1)
641
+ nll_bank = -lb.gather(-1, bidx.clamp_min(0).unsqueeze(-1)).squeeze(-1)
642
+ nll_byte = -self._byte_logprob(h, targets)
643
+ nll = torch.where(in_bank, -log_g + nll_bank, -log_1mg + nll_byte)
644
+ loss = nll.mean()
645
+ if sampled and (step % max(cfg.log_every, 1) == 0):
646
+ nb_full = self._pmix_full_nll(h.detach(), self._bank_index(targets),
647
+ in_bank)
648
+ nll_h = torch.where(in_bank, -log_g.detach() + nb_full,
649
+ -log_1mg.detach() + nll_byte.detach())
650
+ logs["bpb"] = nll_h.mean().item() / 3 / math.log(2) # honest
651
+ logs["bpb_s"] = loss.item() / 3 / math.log(2) # sampled
652
+ else:
653
+ logs["bpb"] = loss.item() / 3 / math.log(2)
654
+ with torch.no_grad(): # branch-conditional currencies
655
+ if in_bank.any():
656
+ logs["bpb_bank_cond"] = nll_bank[in_bank].mean().item() / 3 / math.log(2)
657
+ if (~in_bank).any():
658
+ logs["bpb_byte_cond"] = nll_byte[~in_bank].mean().item() / 3 / math.log(2)
659
+ logs["gate_acc"] = ((torch.sigmoid(g_logit) > 0.5) == in_bank
660
+ ).float().mean().item()
661
+
662
+ # ── auxiliaries ──
663
+ if cfg.write_weight > 0:
664
+ dz, valid = self._write_target(ids)
665
+ pred = F.log_softmax(self.W_write(h), dim=-1)
666
+ kl = F.kl_div(pred, dz, reduction="none").sum(-1)
667
+ wl = kl[:, valid].mean()
668
+ loss = loss + cfg.write_weight * wl
669
+ logs["write_kl"] = wl.item()
670
+
671
+ return loss, logs
672
+
673
+
674
+ # ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
675
+ # Branch-predictive generation (streaming, gauge-triggered forking)
676
+ # ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
677
+ # The CPU metaphor, implemented: per-branch state is the streaming
678
+ # (MΒ±, zΒ±) per layer β€” O(KΒ·d) regardless of context, so speculation is
679
+ # cheap. The model's own multimodality estimate is the branch-predictor
680
+ # confidence bit: fork ONLY where the next-distribution is genuinely
681
+ # multimodal (top-2 probability ratio above `fork_ratio`); run straight
682
+ # through deterministic stretches at full pipeline speed.
683
+ # Position embedding is absolute: generation is capped at cfg.seq_len
684
+ # total positions (documented v1 limitation).
685
+
686
+ def stream_backbone(self, ids: Tensor, states: Optional[List],
687
+ pos_offset: int) -> Tuple[Tensor, List]:
688
+ """One streamed segment (B, S_seg, 3) with per-layer carried states."""
689
+ x = sum(emb(ids[..., i]) for i, emb in enumerate(self.byte_emb))
690
+ x = x + self._pos_slice(ids.shape[1], pos_offset)
691
+ states = states or [None] * len(self.layers)
692
+ new_states: List = []
693
+ for L, st in zip(self.layers, states):
694
+ a, ns = L["attn"].forward_stream(L["norm1"](x), state=st)
695
+ x = x + a
696
+ x = x + L["mlp"](L["norm2"](x))
697
+ new_states.append(ns)
698
+ return self.norm_f(x), new_states
699
+
700
+ @torch.no_grad()
701
+ def next_distribution(self, h_last: Tensor) -> Tuple[Tensor, Tensor, Tensor]:
702
+ """(gate_prob, bank_probs (M,), byte_logprobs (3,256)) for one position.
703
+ h_last: (1, d). Bank scorer follows cfg.bank_scorer."""
704
+ cfg = self.cfg
705
+ g = torch.sigmoid(self.gate(h_last)).squeeze()
706
+ if cfg.bank_scorer == "pmix":
707
+ logits, _ = self._pmix_logits(h_last.unsqueeze(0))
708
+ bank_p = F.softmax(logits.squeeze(0).squeeze(0), dim=-1)
709
+ else:
710
+ pi_p, pi_m = self._pi(h_last)
711
+ k_p, k_m = self._kappa(self.bank)
712
+ bank_p = F.softmax(self._bank_logits(pi_p, pi_m, k_p, k_m
713
+ ).squeeze(0), dim=-1)
714
+ byte_lp = torch.stack([F.log_softmax(head(h_last).squeeze(0), dim=-1)
715
+ for head in self.byte_heads]) # (3,256)
716
+ return g, bank_p, byte_lp
717
+
718
+ @torch.no_grad()
719
+ def generate_tree(self, prompt: bytes, max_new: int = 24,
720
+ beam: int = 8, fork_ratio: float = 0.35,
721
+ fork_width: int = 3, device: str = "cpu") -> List[Dict]:
722
+ """Branch-predictive decoding. Returns the surviving branches as
723
+ [{'text', 'logp', 'forks'}], best first.
724
+ fork_ratio: fork iff p2/p1 > ratio (the confidence bit);
725
+ fork_width: children per fork; beam: global survivor cap."""
726
+ self.eval()
727
+ cfg = self.cfg
728
+ b = prompt[: 3 * (len(prompt) // 3)] or b" "
729
+ ids = torch.frombuffer(bytearray(b), dtype=torch.uint8) \
730
+ .to(torch.long).view(1, -1, 3).to(device)
731
+ h, states = self.stream_backbone(ids, None, 0)
732
+ pos = ids.shape[1]
733
+ Branch = lambda st, h_, lp, txt, forks: \
734
+ {"states": st, "h": h_, "logp": lp, "bytes": txt, "forks": forks}
735
+ branches = [Branch(states, h[:, -1], 0.0, b"", 0)]
736
+
737
+ bank_bytes = self.bank.cpu().numpy().astype("uint8") \
738
+ if self.bank is not None else None
739
+ for step in range(max_new):
740
+ if cfg.pos_mode == "absolute" and pos + 1 > cfg.seq_len:
741
+ break
742
+ nxt: List[Dict] = []
743
+ for br in branches:
744
+ g, bank_p, byte_lp = self.next_distribution(br["h"])
745
+ # mixture distribution over candidate continuations:
746
+ # in-bank candidates carry g*bank_p; the byte tail is
747
+ # summarized by its argmax trigram carrying (1-g)*p_byte
748
+ cand_p, cand_tri = [], []
749
+ if bank_bytes is not None:
750
+ top_p, top_i = bank_p.topk(min(fork_width + 1, len(bank_p)))
751
+ for p, i in zip(top_p.tolist(), top_i.tolist()):
752
+ cand_p.append(g.item() * p)
753
+ cand_tri.append(bytes(bank_bytes[i]))
754
+ by = byte_lp.argmax(-1)
755
+ p_by = float(byte_lp.max(-1).values.sum().exp())
756
+ cand_p.append((1 - g.item()) * p_by)
757
+ cand_tri.append(bytes(by.tolist()))
758
+ order = np.argsort(cand_p)[::-1]
759
+ p1 = cand_p[order[0]]
760
+ p2 = cand_p[order[1]] if len(order) > 1 else 0.0
761
+ take = order[: fork_width] if (p1 > 0 and p2 / max(p1, 1e-12)
762
+ > fork_ratio) else order[:1]
763
+ forked = len(take) > 1
764
+ for oi in take:
765
+ tri = cand_tri[oi]
766
+ t = torch.tensor(list(tri), dtype=torch.long,
767
+ device=device).view(1, 1, 3)
768
+ h2, st2 = self.stream_backbone(
769
+ t, [tuple(s.clone() for s in st)
770
+ for st in br["states"]], pos)
771
+ nxt.append(Branch(st2, h2[:, -1],
772
+ br["logp"] + math.log(max(cand_p[oi], 1e-12)),
773
+ br["bytes"] + tri,
774
+ br["forks"] + int(forked)))
775
+ nxt.sort(key=lambda d: -d["logp"])
776
+ branches = nxt[:beam]
777
+ pos += 1
778
+ return [{"text": (prompt + br["bytes"]).decode("utf-8", errors="replace"),
779
+ "logp": br["logp"], "forks": br["forks"]}
780
+ for br in branches]
781
+
782
+ # ---- the branching gauge ([TAU] inverted) ----
783
+ @torch.no_grad()
784
+ def branching_gauge(self, ids: Tensor, n_baseline: int = 4096) -> Dict:
785
+ h = self.backbone(ids)
786
+ pi_p, pi_m = self._pi(h)
787
+ A = F.normalize(self.codebook, dim=-1)
788
+ conf = ((pi_p - pi_m) @ A).norm(dim=-1).reshape(-1)
789
+ rows = F.normalize(torch.randn(n_baseline, self.cfg.D_addr,
790
+ device=h.device), dim=-1)
791
+ bp, bm = self._address_rows(rows)
792
+ base = ((bp - bm) @ A).norm(dim=-1)
793
+ mu, sd = base.mean(), base.std()
794
+ return {"conf_mean": conf.mean().item(),
795
+ "kernel_invariant": mu.item(),
796
+ "branching_frac": (conf < mu - 2 * sd).float().mean().item()}
797
+
798
+
799
+ # ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
800
+ # Training
801
+ # ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
802
+
803
+ def train_aleph_lm(cfg: AlephLMConfig,
804
+ stream: Optional[TrigramStream] = None) -> Dict:
805
+ torch.manual_seed(cfg.seed)
806
+ dev = torch.device(cfg.device)
807
+ stream = stream or TrigramStream(cfg.corpus_id, cfg.split,
808
+ cfg.max_corpus_bytes, cfg.seed)
809
+ bank = None
810
+ if cfg.head in ("hybrid", "bank", "pointer"):
811
+ if os.path.isfile(str(cfg.bank_source)): # stratified-atlas bank
812
+ d = torch.load(cfg.bank_source, map_location="cpu", weights_only=False)
813
+ bank = d["bank"] if isinstance(d, dict) else d
814
+ print(f"[bank] loaded {len(bank)} trigram candidates "
815
+ f"from {cfg.bank_source}")
816
+ elif cfg.bank_source == "wordnet":
817
+ try:
818
+ bank = build_wordnet_bank(cfg.bank_size)
819
+ print(f"[bank] wordnet char_eng_3gram: {len(bank)} types")
820
+ except Exception as e:
821
+ print(f"[bank] wordnet unavailable ({e}); falling back to corpus")
822
+ if bank is None:
823
+ bank = build_corpus_bank(stream, cfg.bank_size)
824
+ print(f"[bank] corpus top-{len(bank)} trigrams")
825
+
826
+ model = AlephLM(cfg, bank=bank).to(dev)
827
+ n_params = sum(p.numel() for p in model.parameters())
828
+ opt = torch.optim.Adam(model.parameters(), lr=cfg.lr) # pure Adam
829
+ sched = (torch.optim.lr_scheduler.CosineAnnealingLR(
830
+ opt, T_max=cfg.steps, eta_min=cfg.lr * 0.1) if cfg.lr_decay else None)
831
+ alephs = model.aleph_layers()
832
+ for a in alephs:
833
+ a.emit_diversity = cfg.div_weight > 0
834
+
835
+ snapshots: List[Tuple[int, Tensor]] = []
836
+ if cfg.snapshot_codebook:
837
+ snapshots.append((0, alephs[0].export_codebook()))
838
+
839
+ print(f"\n=== AlephLM head={cfg.head} pi={cfg.pi_mode} "
840
+ f"bank={cfg.bank_source if bank is not None else '-'} "
841
+ f"params={n_params:,} ctx={cfg.seq_len} tri "
842
+ f"eff.batch={cfg.batch_size * cfg.accum_steps} dev={dev} ===")
843
+ result: Dict = {"head": cfg.head, "params": n_params}
844
+ t0 = time.time()
845
+
846
+ use_amp = cfg.amp and dev.startswith("cuda")
847
+ if cfg.compile_backbone:
848
+ model.backbone = torch.compile(model.backbone) # tensor-out only
849
+
850
+ for step in range(1, cfg.steps + 1):
851
+ opt.zero_grad(set_to_none=True)
852
+ loss_sum, logs_acc = 0.0, {}
853
+ for _ in range(cfg.accum_steps):
854
+ if cfg.train_mode == "stream":
855
+ ids, targets = stream.sample(
856
+ cfg.batch_size, cfg.seq_len * cfg.segments, dev)
857
+ states = None
858
+ for si in range(cfg.segments):
859
+ sl = slice(si * cfg.seq_len, (si + 1) * cfg.seq_len)
860
+ with torch.autocast(device_type="cuda",
861
+ dtype=torch.bfloat16, enabled=use_amp):
862
+ hseg, states = model.stream_backbone(
863
+ ids[:, sl], states, si * cfg.seq_len)
864
+ loss, logs = model.forward_loss(
865
+ ids[:, sl], targets[:, sl], step=step, h=hseg)
866
+ total = loss
867
+ if cfg.div_weight > 0:
868
+ total = total + cfg.div_weight * sum(
869
+ a.diversity_loss() for a in alephs)
870
+ (total / cfg.accum_steps / cfg.segments).backward()
871
+ states = [tuple(t.detach() for t in st) for st in states]
872
+ loss_sum += loss.item() / cfg.segments
873
+ logs_acc = logs
874
+ continue
875
+ ids, targets = stream.sample(cfg.batch_size, cfg.seq_len, dev)
876
+ with torch.autocast(device_type="cuda",
877
+ dtype=torch.bfloat16, enabled=use_amp):
878
+ loss, logs = model.forward_loss(ids, targets, step=step)
879
+ total = loss
880
+ if cfg.div_weight > 0:
881
+ total = total + cfg.div_weight * sum(
882
+ a.diversity_loss() for a in alephs)
883
+ (total / cfg.accum_steps).backward()
884
+ loss_sum += loss.item()
885
+ logs_acc = logs
886
+ loss_avg = loss_sum / cfg.accum_steps
887
+ gnorm = torch.nn.utils.clip_grad_norm_(
888
+ model.parameters(), max(loss_avg, 1.0))
889
+ opt.step()
890
+ if sched is not None:
891
+ sched.step()
892
+
893
+ if step % cfg.log_every == 0 or step == cfg.steps:
894
+ rate = step * cfg.batch_size * cfg.seq_len * cfg.accum_steps \
895
+ / (time.time() - t0)
896
+ line = (f" step {step:6d} loss {loss_avg:.4f} "
897
+ f"bpb {logs_acc.get('bpb', logs_acc.get('bpb_inbank', float('nan'))):.3f} "
898
+ f"|g| {gnorm:.2f} {rate/1e3:.1f}k tri/s")
899
+ if "bpb_s" in logs_acc:
900
+ line += f" bpbS {logs_acc['bpb_s']:.3f}"
901
+ if "coverage" in logs_acc:
902
+ line += f" cov {logs_acc['coverage']:.0%}"
903
+ if "gate_acc" in logs_acc:
904
+ line += f" gate {logs_acc['gate_acc']:.0%}"
905
+ if "nn_exact" in logs_acc:
906
+ line += (f" bpbE {logs_acc.get('bpb_eval', float('nan')):.3f}"
907
+ f" nn {logs_acc['nn_exact']:.0%}"
908
+ f" ang {logs_acc['ang_err_deg']:.1f}/"
909
+ f"{logs_acc['budget_deg']:.1f}deg"
910
+ f" inBudget {logs_acc['in_budget']:.0%}")
911
+ if "bpb_bank_cond" in logs_acc:
912
+ line += (f" inB {logs_acc['bpb_bank_cond']:.3f}"
913
+ f" outB {logs_acc.get('bpb_byte_cond', float('nan')):.3f}")
914
+ if "mode_spread_deg" in logs_acc:
915
+ line += (f" spread {logs_acc['mode_spread_deg']:.0f}deg"
916
+ f" mixH {logs_acc['mix_entropy']:.2f}")
917
+ if "write_kl" in logs_acc:
918
+ line += f" wKL {logs_acc['write_kl']:.3f}"
919
+ model.eval()
920
+ with torch.no_grad():
921
+ ids_p, _ = stream.sample(min(8, cfg.batch_size), cfg.seq_len, dev)
922
+ st = alephs[0].address_stats(model.backbone(ids_p),
923
+ max_rows=200_000)
924
+ bg = model.branching_gauge(ids_p)
925
+ model.train()
926
+ line += (f" ppl {st['perplexity']:.0f}/{st['max_perplexity']:.0f}"
927
+ f" conf {bg['conf_mean']:.3f}"
928
+ f"/{bg['kernel_invariant']:.3f}"
929
+ f" branch {bg['branching_frac']:.0%}")
930
+ print(line)
931
+ result.update(logs_acc)
932
+ result.update({"loss": loss_avg, "step": step, **bg})
933
+ if cfg.snapshot_codebook:
934
+ snapshots.append((step, alephs[0].export_codebook()))
935
+
936
+ if snapshots:
937
+ traj = [(s, statute(cb)) for s, cb in snapshots]
938
+ result["statute_trajectory"] = traj
939
+ torch.save({"snapshots": snapshots, "statute_trajectory": traj,
940
+ "config": cfg.__dict__}, cfg.snapshot_path)
941
+ d0, d1 = traj[0][1]["deviation"], traj[-1][1]["deviation"]
942
+ print(f"\n[basin] statute: dev {d0:+.4f} -> {d1:+.4f} "
943
+ f"({traj[-1][1]['statute']}); snapshots -> {cfg.snapshot_path}")
944
+ if cfg.checkpoint_path:
945
+ torch.save({"model_state_dict": model.state_dict(),
946
+ "config": cfg.__dict__,
947
+ "bank": model.bank.cpu() if model.bank is not None else None},
948
+ cfg.checkpoint_path)
949
+ print(f"[ckpt] -> {cfg.checkpoint_path}")
950
+ return result
951
+
952
+
953
+ # ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
954
+ # Smoke + activation
955
+ # ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
956
+
957
+ def _smoke():
958
+ print("=" * 70)
959
+ print("AlephLM β€” smoke")
960
+ print("=" * 70)
961
+ rng = np.random.default_rng(0)
962
+ words = [b"the", b"aleph", b"predicts", b"its", b"own", b"future",
963
+ b"through", b"a", b"codebook"]
964
+ path = "/tmp/_alm_corpus.txt"
965
+ with open(path, "wb") as f:
966
+ f.write(b" ".join(words[i] for i in rng.integers(0, 9, 80000)))
967
+
968
+ base = dict(corpus_id=path, max_corpus_bytes=None, steps=25, log_every=25,
969
+ dim=96, n_layers=2, n_heads=4, K=16, seq_len=48, batch_size=8,
970
+ bank_size=256, n_negatives=128, device="cpu",
971
+ checkpoint_path=None, snapshot_path="/tmp/_alm_snaps.pt")
972
+ prior_bpb = 8.0
973
+ for head in ("hybrid", "byte", "bank", "sampled"):
974
+ r = train_aleph_lm(AlephLMConfig(head=head, **base))
975
+ bpb = r.get("bpb", r.get("bpb_inbank", float("nan")))
976
+ assert math.isfinite(r["loss"]), head
977
+ print(f" βœ“ head={head:8s} loss {r['loss']:.3f} bpb {bpb:.2f} "
978
+ f"(uniform prior {prior_bpb:.1f})")
979
+ # pi ablation path + gradient to codebook through the PREDICT/CANDIDATE legs
980
+ cfg = AlephLMConfig(head="hybrid", pi_mode="address", **base)
981
+ stream = TrigramStream(path, max_corpus_bytes=None, seed=0)
982
+ bank = build_corpus_bank(stream, cfg.bank_size)
983
+ m = AlephLM(cfg, bank=bank)
984
+ ids, tg = stream.sample(4, cfg.seq_len, "cpu")
985
+ loss, _ = m.forward_loss(ids, tg)
986
+ loss.backward()
987
+ assert m.codebook.grad is not None and torch.isfinite(m.codebook.grad).all()
988
+ print(f" βœ“ pi_mode='address' ablation runs; codebook grad |{m.codebook.grad.norm():.3f}|")
989
+ print("All smoke tests passed.")
990
+
991
+
992
+ def tier_a_config(bank_source: str, **overrides) -> AlephLMConfig:
993
+ """Tier A of the scaling plan (~25-30M params, one Blackwell, days):
994
+ d=512 x 8 layers, J=16 pmix on a dense bank with sampled softmax,
995
+ TBPTT streaming (4 x 1024 = effective 4096-trigram context at constant
996
+ memory), clamped position (unbounded generation), bf16. Pure Adam,
997
+ standing clip rule β€” nothing exotic enters with scale."""
998
+ base = dict(dim=512, n_layers=8, n_heads=8, K=64, d_point=48,
999
+ head="hybrid", bank_scorer="pmix", n_pointers=16,
1000
+ bank_source=bank_source, bank_softmax="sampled",
1001
+ n_bank_samples=8192, pos_mode="clamp", train_mode="stream",
1002
+ segments=4, seq_len=1024, batch_size=16, accum_steps=4,
1003
+ lr=5e-4, steps=50_000, amp=True, log_every=250)
1004
+ base.update(overrides)
1005
+ return AlephLMConfig(**base)
1006
+
1007
+
1008
+ if __name__ == "__main__":
1009
+ import argparse
1010
+ ap = argparse.ArgumentParser(description="AlephLM β€” prediction through the codebook")
1011
+ ap.add_argument("--smoke-only", action="store_true")
1012
+ ap.add_argument("--head", default="hybrid",
1013
+ choices=["hybrid", "byte", "bank", "sampled", "pointer"])
1014
+ ap.add_argument("--bank", default="corpus", choices=["corpus", "wordnet"])
1015
+ ap.add_argument("--pi", default="free", choices=["free", "address"])
1016
+ ap.add_argument("--scorer", default="kernel", choices=["kernel", "pmix"])
1017
+ ap.add_argument("--pointers", type=int, default=4)
1018
+ ap.add_argument("--steps", type=int, default=10_000)
1019
+ ap.add_argument("--corpus-mb", type=int, default=100)
1020
+ ap.add_argument("--device",
1021
+ default="cuda" if torch.cuda.is_available() else "cpu")
1022
+ args, _unknown = ap.parse_known_args()
1023
+ if args.smoke_only:
1024
+ _smoke()
1025
+ else:
1026
+ cfg = AlephLMConfig(head=args.head, bank_source=args.bank,
1027
+ pi_mode=args.pi, steps=args.steps,
1028
+ bank_scorer=args.scorer, n_pointers=args.pointers,
1029
+ max_corpus_bytes=args.corpus_mb * 1_000_000,
1030
+ device=args.device)
1031
+ train_aleph_lm(cfg)