Quazim0t0 commited on
Commit
790650c
·
verified ·
1 Parent(s): 6427525

Add Mycel-LM 79M: architecture, tokenizer, training scripts, weights-only checkpoints

Browse files
Files changed (25) hide show
  1. README.md +317 -0
  2. build_fractal_table.py +106 -0
  3. chat_sft.py +129 -0
  4. chkpt/quazimoto.pt +3 -0
  5. chkpt/quazimoto_sft.pt +3 -0
  6. distill_uld.py +194 -0
  7. family.py +446 -0
  8. fractal.py +79 -0
  9. fractal_phase.pt +3 -0
  10. generate.py +326 -0
  11. healthcheck.py +258 -0
  12. instrument.py +84 -0
  13. model.py +773 -0
  14. mycel.py +178 -0
  15. opd_teacher.py +79 -0
  16. requirements.txt +11 -0
  17. special_tokens.py +85 -0
  18. spike_tokenizer.py +124 -0
  19. tokenizer.json +0 -0
  20. train.bat +24 -0
  21. train.py +297 -0
  22. train_opd.py +290 -0
  23. train_sft.bat +16 -0
  24. train_sft.py +235 -0
  25. visualize.py +287 -0
README.md ADDED
@@ -0,0 +1,317 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ library_name: pytorch
4
+ pipeline_tag: text-generation
5
+ tags:
6
+ - custom-architecture
7
+ - kuramoto
8
+ - oscillator
9
+ - mycelium
10
+ - neighbour-sensing
11
+ - experimental
12
+ - research
13
+ - spikewhale-tokenizer
14
+ ---
15
+
16
+ # Mycel-LM (79M)
17
+
18
+ **Mycel-LM** is a **79.2M-parameter** research language model whose channel-mixing block is
19
+ **not** an MLP. It is a differentiable **Neighbour-Sensing fungal-colony-growth** model:
20
+ each token is expanded into a colony of *hyphal tips* that grow in a bounded latent
21
+ region, sense a shared density field, and steer their own growth — the "MLP" is replaced
22
+ by a few differentiable steps of colony growth, read back out into the hidden state.
23
+
24
+ It is part of a family of models that ask a single question: **can the generalizing
25
+ ability of a transformer be carried by an unusual, self-organizing dynamical system in
26
+ place of the feed-forward block?** Mycel-LM keeps the family's tokenizer, traits, and
27
+ data fixed and swaps *only* the mixer, so it is a controlled experiment against the
28
+ sibling **Quazimoto** models (whose mixer is a bank of coupled Kuramoto oscillators).
29
+
30
+ > ⚠️ **Research artifact, not a product.** At ~79M parameters it is fluent but small: it
31
+ > models the *shape* of language well and generates coherent, grammatical text, but it is
32
+ > **not factual** and will confidently hallucinate. See [Limitations](#limitations).
33
+
34
+ ---
35
+
36
+ ## Table of contents
37
+
38
+ - [Highlights](#highlights)
39
+ - [Architecture](#architecture)
40
+ - [Repository layout](#repository-layout)
41
+ - [Install](#install)
42
+ - [Quickstart](#quickstart)
43
+ - [Command-line usage](#command-line-usage)
44
+ - [Live visualizer & Space](#live-visualizer--space)
45
+ - [Training from scratch](#training-from-scratch)
46
+ - [Fine-tuning (SFT)](#fine-tuning-sft)
47
+ - [Checkpoints](#checkpoints)
48
+ - [Limitations](#limitations)
49
+ - [Citation / basis](#citation--basis)
50
+
51
+ ---
52
+
53
+ ## Highlights
54
+
55
+ - **Novel mixer.** The per-layer feed-forward block is replaced by a **MycelBlock** — a
56
+ differentiable simulation of fungal colony growth (Neighbour-Sensing).
57
+ - **Self-describing checkpoints.** Each `.pt` embeds a `family_config` recording the exact
58
+ geometry, so `generate.py` / `healthcheck.py` / `visualize.py` rebuild the model with no
59
+ external config.
60
+ - **KV cache.** Incremental decoding is wired through the whole stack (attention presents
61
+ are threaded per layer); `generate()` prefills the prompt once and decodes one token per
62
+ forward.
63
+ - **Self-speculative decoding.** Four MTP draft heads propose the next tokens and the main
64
+ head verifies them in one parallel forward — bit-identical to greedy, just fewer forwards.
65
+ - **Live 3-D visualizer.** Watch the colony grow token-by-token as a Three.js filament web.
66
+
67
+ ---
68
+
69
+ ## Architecture
70
+
71
+ Standard causal Transformer backbone (token-mixing = attention, tied LM head), with the
72
+ per-layer feed-forward network replaced by a **MycelBlock**.
73
+
74
+ ### The MycelBlock (the novel part)
75
+
76
+ Based on the Meškauskas / Fricker / Moore (2004) **Neighbour-Sensing** model of fungal
77
+ colony growth:
78
+
79
+ 1. The hidden state projects to **N = 96 hyphal tips** per token, each with a **position**
80
+ in a bounded 3-D latent region and a **growth vector**.
81
+ 2. A few differentiable **growth steps** run: each tip senses the local **density field**,
82
+ steers *away* from the colony's own density (negative autotropism) with persistence,
83
+ moves, and is **re-clamped** into the bounded region (the colony can't grow unbounded).
84
+ 3. The final `[position, growth-vector, sensed-density]` of every tip is read out back into
85
+ the hidden state, behind a family gate.
86
+
87
+ The density field is evaluated against **16 learnable field centres** (a low-rank sample of
88
+ the field) so cost is **O(N·F)** per step, not O(N²) — the same mean-field trick that keeps
89
+ the sibling oscillator block cheap. Health-checking a trained checkpoint shows the tropism
90
+ parameter converges **strongly negative** across layers, i.e. the model genuinely learns the
91
+ grow-away-from-density behaviour rather than leaving it at init.
92
+
93
+ **Trait stations** (`MycelStations`): tiny memory specialists sit at fixed anchor positions
94
+ in the colony. A tip interacts with a station by **proximity** — which is *emergent* from
95
+ where the tip grew — so which tips use which trait "comes to be" during growth rather than
96
+ being assigned to a fixed index. The stations hold test-time-writable input/output stores
97
+ that act as an addressable context memory at inference.
98
+
99
+ ### Attention
100
+
101
+ Family attention ported from the Quazimoto v2 stack:
102
+ - **MLA** low-rank Q/O projections
103
+ - **Partial RoPE** (nope + rope split), **QK-Norm**, **GQA** (4 KV heads)
104
+ - optional DERF (erf attention) and XSA (value-subspace removal) — off in this checkpoint
105
+ - **KV cache** for incremental decoding (per-layer `(k, v)` presents threaded through the stack)
106
+
107
+ ### Opt-in family traits (all live in this checkpoint)
108
+
109
+ | Trait | Role |
110
+ |---|---|
111
+ | **HRM** | iterative gated hidden-state refinement (random init state, gates open) |
112
+ | **MoE** | SwiGLU mixture (4 routed + 1 shared, top-2) refining the trunk |
113
+ | **MTP** (×4) | multi-token-prediction draft heads → enables self-speculative decoding |
114
+ | **JEPA** | representation-prediction aux loss (train-only; never runs at inference) |
115
+ | **Ring Specialists** (7/ring) | the trait stations described above |
116
+ | **Fractal Phase Seed** | seeds tip positions from each token's Mandelbrot orbit angles (gated) |
117
+
118
+ ### Config (this checkpoint)
119
+
120
+ | | |
121
+ |---|---|
122
+ | params | **79.2M** |
123
+ | layers | 10 |
124
+ | d_model | 768 |
125
+ | heads | 12 (4 KV) |
126
+ | vocab | 16512 (SpikeWhale byte-merge) |
127
+ | block size | 2048 |
128
+ | tips / token | 96, in a 3-D bounded colony |
129
+ | field centres | 16 · growth steps 3 · stations 16 |
130
+
131
+ The checkpoint is **self-describing**: `family_config` inside the `.pt` records the exact
132
+ geometry so the model rebuilds itself on load.
133
+
134
+ ---
135
+
136
+ ## Repository layout
137
+
138
+ ```
139
+ model.py QuazimotoLM + QuazimotoConfig — the transformer backbone, attention,
140
+ KV cache, traits (HRM/MoE/MTP/JEPA), generate() and forward_drafts()
141
+ mycel.py MycelBlock (Neighbour-Sensing growth mixer) + MycelStations
142
+ family.py shared family layers (MoE, HRM, specialists, norms, ...)
143
+ fractal.py hierarchical Mandelbrot phase seeding (FractalSeed trait)
144
+ instrument.py zero-cost capture hooks the visualizer reads from
145
+ special_tokens.py ChatML / control-token definitions
146
+ spike_tokenizer.py SpikeWhale byte-merge tokenizer (subclasses PreTrainedTokenizer)
147
+ tokenizer.json the tokenizer vocab / merges (vocab 16,512)
148
+ fractal_phase.pt precomputed hierarchical Mandelbrot phase table (regenerable)
149
+
150
+ generate.py inference harness — KV cache + self-speculative decoding + sampling
151
+ healthcheck.py per-layer weight / gate / PPL diagnostics for a checkpoint
152
+ visualize.py builds the 3-D colony dashboard (viz.html) from a generation
153
+
154
+ train.py pretraining entry point (streamed multi-corpus blend)
155
+ train_sft.py supervised fine-tuning (ChatML, assistant-only loss masking)
156
+ chat_sft.py chat-format rendering / loss masking helpers used by SFT
157
+ train_opd.py OPD (on-policy distillation) training loop
158
+ distill_uld.py universal-logit-distillation utilities
159
+ opd_teacher.py teacher wrapper for distillation
160
+ build_fractal_table.py regenerates fractal_phase.pt
161
+ train.bat / train_sft.bat Windows convenience launchers
162
+
163
+ chkpt/quazimoto.pt pretraining checkpoint (step 149,000)
164
+ chkpt/quazimoto_sft.pt SFT checkpoint (step 4,000, ChatML)
165
+ ```
166
+
167
+ > **Note:** the Modal cloud launchers (`modal_train.py`, `modal_sft.py`) are intentionally
168
+ > **not** part of this package. The scripts above run locally on CPU or a single GPU.
169
+
170
+ ---
171
+
172
+ ## Install
173
+
174
+ ```bash
175
+ pip install -r requirements.txt
176
+ ```
177
+
178
+ Requirements are minimal: `torch`, `numpy`, `transformers` (the tokenizer subclasses
179
+ `PreTrainedTokenizer`). Training additionally uses `datasets` and `huggingface_hub`.
180
+ Everything below runs on **CPU** (slow but functional) or a single GPU.
181
+
182
+ ---
183
+
184
+ ## Quickstart
185
+
186
+ ```python
187
+ import torch
188
+ from model import QuazimotoLM, QuazimotoConfig
189
+ from spike_tokenizer import SpikeTokenizer
190
+
191
+ ck = torch.load("chkpt/quazimoto.pt", map_location="cpu", weights_only=False)
192
+ cfg = QuazimotoConfig(**ck["family_config"]) # self-describing
193
+ model = QuazimotoLM(cfg); model.load_state_dict(ck["model"], strict=False); model.eval()
194
+ tok = SpikeTokenizer(vocab_file="tokenizer.json")
195
+
196
+ ids = torch.tensor([tok.encode("The mycelium spreads through the soil", add_special_tokens=False)])
197
+ out = model.generate(ids, n_new=80, temperature=0.8, top_k=40) # KV cache on by default
198
+ print(tok.decode(out[0].tolist(), skip_special_tokens=True))
199
+ ```
200
+
201
+ For a **chat** turn, wrap the prompt in ChatML and stop on `<|im_end|>` (the SFT checkpoint
202
+ was trained on this framing):
203
+
204
+ ```python
205
+ prompt = "<|im_start|><|user|>\nWhat is mycelium?<|im_end|>\n<|im_start|><|assistant|>\n"
206
+ ids = torch.tensor([tok.encode(prompt, add_special_tokens=False)])
207
+ out = model.generate(ids, n_new=120, temperature=0.7, top_k=40)
208
+ ```
209
+
210
+ ---
211
+
212
+ ## Command-line usage
213
+
214
+ ```bash
215
+ # plain completion (KV cache on by default)
216
+ python generate.py --ckpt chkpt/quazimoto.pt --prompt "In the beginning" --max_new_tokens 80
217
+
218
+ # chat turn (ChatML framing + stop on <|im_end|>)
219
+ python generate.py --ckpt chkpt/quazimoto_sft.pt --chat --prompt "Hello, who are you?"
220
+
221
+ # interactive REPL
222
+ python generate.py --ckpt chkpt/quazimoto_sft.pt --interactive
223
+
224
+ # self-speculative decoding (MTP heads draft, main head verifies; report acceptance)
225
+ python generate.py --ckpt chkpt/quazimoto.pt --speculative --spec_stats
226
+
227
+ # disable the KV cache (full recompute each step — for comparison)
228
+ python generate.py --ckpt chkpt/quazimoto.pt --no_cache
229
+
230
+ # per-layer diagnostics (weights / gates / PPL)
231
+ python healthcheck.py --ckpt chkpt/quazimoto.pt
232
+ ```
233
+
234
+ Sampling knobs: `--temperature`, `--top_k`, `--top_p`, `--repetition_penalty`, `--seed`.
235
+
236
+ ---
237
+
238
+ ## Live visualizer & Space
239
+
240
+ `visualize.py` renders the colony growing in 3-D as the model generates, token by token —
241
+ hyphal tips linked into a filament web, coloured by local density, with the trait stations
242
+ shown as orange wire-spheres. It writes a self-contained `viz.html` (Three.js from a CDN):
243
+
244
+ ```bash
245
+ python visualize.py --ckpt chkpt/quazimoto_sft.pt --prompt "the mycelium spreads" --tokens 50
246
+ ```
247
+
248
+ A companion **Hugging Face Space (`Mycel-LM v1`)** wraps the same architecture in an
249
+ interactive chat — KV-cache decoding drives the reply while the 3-D colony visualizer
250
+ animates the growth for the generated tokens.
251
+
252
+ ---
253
+
254
+ ## Training from scratch
255
+
256
+ ```bash
257
+ python train.py --device cuda --steps 160000 --batch 12 --block 2048 --amp \
258
+ --use-hrm --use-moe --use-mtp --use-jepa --use-ring-specialists --use-fractal-phase-seed \
259
+ --stream --math-frac 0.25 --out chkpt/quazimoto.pt --ckpt-every 500 --resume
260
+ ```
261
+
262
+ - **Tokenizer:** SpikeWhale byte-merge, vocab 16,512. (Byte-merge perplexity is
263
+ tokenizer-inflated; **bits/byte** is the honest metric.)
264
+ - **Pretraining blend:** 35% Ultra-FineWeb-L3 / 25% FineWeb-Edu / 25% FineMath /
265
+ 15% Quazim0t0/PretrainNew, streamed. Streamed datasets are pulled with `datasets`;
266
+ gated corpora need `huggingface-cli login`.
267
+ - `--resume` continues from the checkpoint at `--out`. The growth loop is activation-heavy,
268
+ so keep the batch modest; `--amp` gives a bf16 speedup on GPU.
269
+
270
+ Pass `--help` to `train.py` for the full trait / optimiser / schedule surface.
271
+
272
+ ## Fine-tuning (SFT)
273
+
274
+ ```bash
275
+ python train_sft.py --init chkpt/quazimoto.pt --out chkpt/quazimoto_sft.pt \
276
+ --steps 4000 --batch 8 --block 2048 --amp
277
+ ```
278
+
279
+ - Renders a chat mix in **ChatML** with **assistant-only loss masking** (`chat_sft.py`).
280
+ - **SFT blend:** ultrachat_200k_sft + ultrafeedback-sft + UltraData-SFT-2605/Knowledge +
281
+ OpenThoughts2-1M-ShortThink.
282
+ - The bundled SFT checkpoint is only **4k steps** — the chat format transferred but the
283
+ model is still shallow.
284
+
285
+ > The distributed checkpoints carry **weights only** (optimizer state stripped to keep the
286
+ > download small). Fine-tuning starts a fresh optimizer from them, which is the normal path;
287
+ > only exact *resumption* of the original pretraining run would need the optimizer state.
288
+
289
+ ---
290
+
291
+ ## Checkpoints
292
+
293
+ - `chkpt/quazimoto.pt` — **pretraining** checkpoint, step **149,000**
294
+ - `chkpt/quazimoto_sft.pt` — **SFT** checkpoint, step **4,000** (ChatML, early)
295
+
296
+ Both embed `family_config` (self-describing) and load with `strict=False` so future trait
297
+ additions stay backward-compatible.
298
+
299
+ ---
300
+
301
+ ## Limitations
302
+
303
+ - **Not factual.** Small-model behaviour: fluent and grammatical, but it invents facts
304
+ ("the capital of France is the largest and most important part of the world").
305
+ - **SFT is early** (4k steps) — answers follow the chat format but hallucinate.
306
+ - **No safety tuning.** No RLHF/guardrails; do not deploy in user-facing settings.
307
+ - **Custom architecture** — cannot be loaded with `AutoModel`; use the bundled `model.py`.
308
+ - This is an **experiment in architecture**, released to study whether a self-organizing
309
+ growth process can carry a transformer's generalization. Treat outputs accordingly.
310
+
311
+ ## Citation / basis
312
+
313
+ Neighbour-Sensing model of hyphal growth: Meškauskas, Fricker & Moore (2004),
314
+ *Simulating colonial growth of fungi with the Neighbour-Sensing model of hyphal growth*,
315
+ Mycological Research 108(11).
316
+
317
+ *License: Apache-2.0.*
build_fractal_table.py ADDED
@@ -0,0 +1,106 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ build_fractal_table.py -- precompute the HIERARCHICAL (tokenizer-aware) Mandelbrot
3
+ phase-seed table for Quazimoto-LM. CPU-only, run once; the model loads the result.
4
+
5
+ Why offline: the hierarchical map needs the tokenizer (each token's byte string),
6
+ which the model deliberately doesn't carry. So we build the [vocab, n_osc] table
7
+ here and save it; model.py loads it when --use-fractal-phase-seed is set.
8
+
9
+ Hierarchy via an IFS address (this is exactly how fractal self-similarity encodes
10
+ prefixes):
11
+ * each of the 256 byte VALUES gets a fixed anchor point in the Mandelbrot region
12
+ (Halton spread, so anchors are evenly distributed and deterministic);
13
+ * a token's c starts at its FIRST byte's anchor and is contracted toward each
14
+ subsequent byte's anchor (p <- (1-w)*p + w*anchor[b]).
15
+ Tokens sharing a prefix therefore land in nested, self-similar neighborhoods
16
+ (prefix dominates; later bytes refine at shrinking scale). Specials (ids 0-3 and
17
+ the appended <|...|> set) get their own deterministic anchors away from the bytes.
18
+ Then phase_k = angle(z_k) of z <- z^2 + c, identical to the flat builder.
19
+
20
+ Usage:
21
+ python build_fractal_table.py # default n_osc from QuazimotoConfig
22
+ python build_fractal_table.py --n-osc 508 --out fractal_phase.pt
23
+ """
24
+
25
+ import argparse
26
+ import os
27
+ import sys
28
+
29
+ import torch
30
+
31
+ from fractal import _halton, phases_from_c
32
+ from model import QuazimotoConfig
33
+
34
+ PKG_DIR = os.path.dirname(os.path.abspath(__file__))
35
+
36
+
37
+ def load_tokenizer(tok_dir):
38
+ sys.path.insert(0, tok_dir)
39
+ from spike_tokenizer import SpikeTokenizer
40
+ return SpikeTokenizer(vocab_file=os.path.join(tok_dir, "tokenizer.json"))
41
+
42
+
43
+ def byte_anchors(region):
44
+ """256 fixed anchor points (one per byte value) spread by Halton over region."""
45
+ x0, x1, y0, y1 = region
46
+ ax = [x0 + _halton(b + 1, 2) * (x1 - x0) for b in range(256)]
47
+ ay = [y0 + _halton(b + 1, 3) * (y1 - y0) for b in range(256)]
48
+ return ax, ay
49
+
50
+
51
+ def build(tok, n_osc, region=(-0.8, -0.2, 0.4, 0.9), contract=0.15):
52
+ # region: an off-origin band near the Mandelbrot boundary so the early orbit
53
+ # angles vary SMOOTHLY with c (angle doesn't wrap wildly as it does across 0)
54
+ # -> prefix locality survives into the low oscillators. contract: small, so the
55
+ # FIRST byte dominates the address and shared-prefix tokens cluster tightly.
56
+ V = tok.vocab_size
57
+ ax, ay = byte_anchors(region)
58
+ specials = set(getattr(tok, "_extra_specials", []))
59
+ ids2tok = tok._ids_to_tokens
60
+ cr = torch.empty(V); ci = torch.empty(V)
61
+ n_byte = n_merge = n_special = 0
62
+ for tid in range(V):
63
+ s = ids2tok.get(tid, "")
64
+ is_special = tid < 4 or s in specials
65
+ if is_special or not s:
66
+ # own corner, away from the byte anchors (different Halton stream)
67
+ cr[tid] = region[0] + _halton(tid + 1, 5) * (region[1] - region[0])
68
+ ci[tid] = region[2] + _halton(tid + 1, 7) * (region[3] - region[2])
69
+ n_special += 1
70
+ continue
71
+ byts = [ord(ch) & 0xFF for ch in s] # latin-1 byte values
72
+ px, py = ax[byts[0]], ay[byts[0]] # first byte dominates
73
+ for b in byts[1:]: # contract toward each later byte
74
+ px = (1 - contract) * px + contract * ax[b]
75
+ py = (1 - contract) * py + contract * ay[b]
76
+ cr[tid], ci[tid] = px, py
77
+ if len(byts) == 1:
78
+ n_byte += 1
79
+ else:
80
+ n_merge += 1
81
+ print(f"mapped {V} tokens -> c ({n_byte} single-byte, {n_merge} merges, {n_special} special)")
82
+ phases = phases_from_c(torch.complex(cr, ci), n_osc)
83
+ return phases
84
+
85
+
86
+ def main():
87
+ p = argparse.ArgumentParser(description="Build hierarchical Mandelbrot phase table (CPU)")
88
+ p.add_argument("--tok_dir", default=PKG_DIR)
89
+ p.add_argument("--n-osc", type=int, default=QuazimotoConfig().n_osc,
90
+ help="number of oscillators (sum of ring_sizes); default from config")
91
+ p.add_argument("--contract", type=float, default=0.15, help="IFS contraction toward later bytes (small = first byte dominates)")
92
+ p.add_argument("--out", default=os.path.join(PKG_DIR, "fractal_phase.pt"))
93
+ args = p.parse_args()
94
+
95
+ torch.set_default_device("cpu")
96
+ tok = load_tokenizer(args.tok_dir)
97
+ print(f"tokenizer vocab {tok.vocab_size} | building hierarchical table, n_osc={args.n_osc}")
98
+ phases = build(tok, args.n_osc, contract=args.contract).cpu()
99
+ torch.save({"phases": phases, "vocab_size": tok.vocab_size, "n_osc": args.n_osc,
100
+ "mode": "hierarchical", "contract": args.contract}, args.out)
101
+ print(f"saved {args.out} shape {tuple(phases.shape)} "
102
+ f"({os.path.getsize(args.out)//1024} KB) | range [{phases.min():.3f}, {phases.max():.3f}]")
103
+
104
+
105
+ if __name__ == "__main__":
106
+ main()
chat_sft.py ADDED
@@ -0,0 +1,129 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ chat_sft.py -- interactive multi-turn chat REPL for an SFT'd Quazimoto-LM.
3
+
4
+ Renders the running conversation in the student's ChatML and generates the
5
+ assistant reply. Two decode paths:
6
+
7
+ * SPECULATIVE (default) -- DeepSpec/DSpark-style self-speculative decoding: the
8
+ MTP heads draft the next few tokens, the main head verifies them in one pass,
9
+ accepting the longest correct prefix (see model.forward_drafts / generate.py
10
+ generate_speculative). Greedy, so output is deterministic.
11
+ * SAMPLED (--temperature > 0) -- the KV-cache sampler (generate.py generate):
12
+ the conversation is prefilled once, then each new token is a single-position
13
+ cached forward. Supports temperature / top-k / top-p.
14
+
15
+ Both reuse the model's KV cache; speculative additionally drafts multiple tokens
16
+ per verify. Usage:
17
+ python chat_sft.py # auto-find newest chkpt, speculative
18
+ python chat_sft.py --temperature 0.7 # sampled instead
19
+ python chat_sft.py --ckpt chkpt/quazimoto_sft.pt --system "You are concise."
20
+ """
21
+ import argparse, os, sys
22
+ import torch
23
+
24
+ from model import QuazimotoLM, QuazimotoConfig
25
+ from generate import generate, generate_speculative, resolve_stop_ids
26
+
27
+ PKG_DIR = os.path.dirname(os.path.abspath(__file__))
28
+ ROLE_TOK = {"system": "<|system|>", "user": "<|user|>", "assistant": "<|assistant|>"}
29
+
30
+
31
+ def find_ckpt(path):
32
+ if path and os.path.isfile(path):
33
+ return path
34
+ folder = (os.path.dirname(path) if path else os.path.join(PKG_DIR, "chkpt")) or "."
35
+ if not os.path.isdir(folder):
36
+ return None
37
+ # prefer an SFT checkpoint, else newest *.pt
38
+ pts = [os.path.join(folder, f) for f in os.listdir(folder) if f.endswith(".pt")]
39
+ sft = [p for p in pts if "sft" in os.path.basename(p).lower()]
40
+ pool = sft or pts
41
+ return max(pool, key=os.path.getmtime) if pool else None
42
+
43
+
44
+ def load_tokenizer(tok_dir):
45
+ sys.path.insert(0, tok_dir)
46
+ from spike_tokenizer import SpikeTokenizer
47
+ return SpikeTokenizer(vocab_file=os.path.join(tok_dir, "tokenizer.json"))
48
+
49
+
50
+ def render(history, tok):
51
+ """history: list of {role,content} -> token ids ending with the assistant header
52
+ so the model continues the assistant turn."""
53
+ parts = [f"<|im_start|>{ROLE_TOK.get(m['role'],'<|user|>')}\n{m['content']}<|im_end|>\n"
54
+ for m in history]
55
+ parts.append("<|im_start|><|assistant|>\n")
56
+ return tok.encode("".join(parts), add_special_tokens=False)
57
+
58
+
59
+ def main():
60
+ p = argparse.ArgumentParser(description="Quazimoto-LM SFT chat (KV cache + speculative)")
61
+ p.add_argument("--ckpt", default="", help="SFT checkpoint (default: newest *sft*.pt in ./chkpt)")
62
+ p.add_argument("--tok_dir", default=PKG_DIR)
63
+ p.add_argument("--system", default="", help="optional system prompt")
64
+ p.add_argument("--max_new_tokens", type=int, default=256)
65
+ p.add_argument("--temperature", type=float, default=0.0, help="0 = speculative/greedy; >0 = sampled")
66
+ p.add_argument("--top_k", type=int, default=40)
67
+ p.add_argument("--top_p", type=float, default=0.95)
68
+ p.add_argument("--repetition_penalty", type=float, default=1.1)
69
+ p.add_argument("--no_speculative", action="store_true", help="force the cached sampler even at temp 0")
70
+ p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
71
+ args = p.parse_args()
72
+
73
+ path = find_ckpt(args.ckpt)
74
+ if path is None:
75
+ print("No checkpoint found; pass --ckpt or train one."); return
76
+ ckpt = torch.load(path, map_location=args.device, weights_only=False)
77
+ cfg = QuazimotoConfig(**ckpt["family_config"])
78
+ model = QuazimotoLM(cfg); model.load_state_dict(ckpt["model"], strict=False)
79
+ model.to(args.device).eval()
80
+ tok = load_tokenizer(args.tok_dir)
81
+ stop = resolve_stop_ids(tok, chat=True) # <|im_end|>, <eos>, <|endoftext|>
82
+ imend = tok.get_vocab().get("<|im_end|>")
83
+
84
+ spec = (args.temperature <= 1e-4 and not args.no_speculative and model.mtp_heads is not None)
85
+ mode = "speculative (DSpark draft+verify)" if spec else f"sampled (temp {args.temperature}, KV cache)"
86
+ print(f"loaded {os.path.basename(path)} (step {ckpt.get('step')}) | decode: {mode}")
87
+ print("type your message; /reset clears history, /quit exits.\n")
88
+
89
+ history = []
90
+ if args.system:
91
+ history.append({"role": "system", "content": args.system})
92
+
93
+ while True:
94
+ try:
95
+ user = input("you> ").strip()
96
+ except (EOFError, KeyboardInterrupt):
97
+ print(); break
98
+ if not user:
99
+ continue
100
+ if user in ("/quit", "/exit"):
101
+ break
102
+ if user == "/reset":
103
+ history = ([{"role": "system", "content": args.system}] if args.system else [])
104
+ print("(history cleared)\n"); continue
105
+
106
+ history.append({"role": "user", "content": user})
107
+ ids = render(history, tok)
108
+ x = torch.tensor([ids], device=args.device)
109
+ with torch.no_grad():
110
+ if spec:
111
+ out = generate_speculative(model, cfg, x, args.max_new_tokens, stop_ids=stop)
112
+ else:
113
+ out = generate(model, cfg, x, args.max_new_tokens, temperature=max(args.temperature, 1e-6),
114
+ top_k=args.top_k or None, top_p=args.top_p,
115
+ repetition_penalty=args.repetition_penalty, stop_ids=stop,
116
+ device=args.device, use_cache=True)
117
+ gen = out[0, len(ids):].tolist()
118
+ if imend in gen: # trim at the assistant turn end
119
+ gen = gen[:gen.index(imend)]
120
+ reply = tok.decode(gen, skip_special_tokens=True).strip()
121
+ try:
122
+ print(f"bot> {reply}\n")
123
+ except UnicodeEncodeError:
124
+ print("bot> " + reply.encode("ascii", "replace").decode() + "\n")
125
+ history.append({"role": "assistant", "content": reply})
126
+
127
+
128
+ if __name__ == "__main__":
129
+ main()
chkpt/quazimoto.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:24688695d1d6540b31958fb143e2b1a34ab726a2bda4f53b79aa74f0c54b3d71
3
+ size 317207627
chkpt/quazimoto_sft.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3f5e88339e4d88140884d360e58952a9c6f867448d4a5608fbf98cd34cda9815
3
+ size 317209099
distill_uld.py ADDED
@@ -0,0 +1,194 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ distill_uld.py -- OFF-policy cross-tokenizer distillation of MobileLLM-R1-140M-base
3
+ (teacher) into Quazimoto-LM (student), using the ULD recipe from distill_mark2.py:
4
+
5
+ loss = CE(student, real tokens) + ULD_W * ULD(student_logits, teacher_logits)
6
+
7
+ ULD (Universal Logit Distillation) is what makes this tokenizer-agnostic: at each
8
+ position it softmaxes both models, SORTS each distribution descending, keeps the
9
+ top-K, and takes an L1 between the sorted vectors -- comparing the *shape* of the
10
+ distributions (how peaked vs. flat) while discarding token identity. So no vocab
11
+ alignment is needed, and unlike the on-policy OPD script there is NO generation
12
+ (teacher-forced), so it runs at normal-training speed.
13
+
14
+ Differences from the HF student in distill_mark2.py, handled here:
15
+ * QuazimotoLM does NOT shift targets internally -> we build next-token targets.
16
+ * its CE uses ignore_index=-1 (not -100).
17
+ * right-padding is safe (causal attention; real tokens never attend the pad tail).
18
+ * we pass targets=None and compute CE manually, so the AR aux heads (MTP/JEPA)
19
+ don't run during distillation.
20
+
21
+ Run on a rig with transformers + the teacher weights:
22
+ python distill_uld.py --student-ckpt chkpt/quazimoto.pt --device cuda --stream
23
+ """
24
+
25
+ import argparse, math, os, time
26
+ import torch
27
+ import torch.nn.functional as F
28
+
29
+ from model import QuazimotoLM, QuazimotoConfig
30
+ import train as T
31
+
32
+ PKG_DIR = os.path.dirname(os.path.abspath(__file__))
33
+
34
+
35
+ def uld_loss(s_logits, t_logits, uld_n, top_k):
36
+ """Sorted-top-K L1 between student and teacher next-token distributions,
37
+ averaged over valid positions. s_logits/t_logits: [B, S, V*] (each its own V)."""
38
+ total_L, total_n = 0.0, 0
39
+ for j, n in enumerate(uld_n):
40
+ if n <= 0:
41
+ continue
42
+ sp = F.softmax(s_logits[j, :n].float(), -1).sort(-1, descending=True).values[:, :top_k]
43
+ tp = F.softmax(t_logits[j, :n].float(), -1).sort(-1, descending=True).values[:, :top_k]
44
+ k = min(sp.size(-1), tp.size(-1)) # vocab could be < top_k
45
+ total_L = total_L + (sp[:, :k] - tp[:, :k]).abs().sum()
46
+ total_n += n
47
+ return total_L / max(total_n, 1)
48
+
49
+
50
+ def main():
51
+ p = argparse.ArgumentParser()
52
+ p.add_argument("--tok-dir", default=PKG_DIR)
53
+ p.add_argument("--student-ckpt", default=os.path.join(PKG_DIR, "chkpt", "quazimoto.pt"))
54
+ p.add_argument("--teacher", default="facebook/MobileLLM-R1-140M-base")
55
+ p.add_argument("--steps", type=int, default=6000)
56
+ p.add_argument("--seq", type=int, default=1024)
57
+ p.add_argument("--micro", type=int, default=2, help="micro-batch (sequences per fwd)")
58
+ p.add_argument("--accum", type=int, default=8, help="grad-accum micro-steps")
59
+ p.add_argument("--lr", type=float, default=1.5e-4)
60
+ p.add_argument("--warmup", type=int, default=100)
61
+ p.add_argument("--min-lr-frac", type=float, default=0.1)
62
+ p.add_argument("--uld-w", type=float, default=0.5, help="ULD loss weight")
63
+ p.add_argument("--top-k", type=int, default=128, help="ULD sorted top-K")
64
+ p.add_argument("--weight-decay", type=float, default=0.01)
65
+ p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
66
+ p.add_argument("--teacher-device", default="")
67
+ p.add_argument("--amp", action="store_true", help="bf16 autocast on the student forward")
68
+ p.add_argument("--out", default=os.path.join(PKG_DIR, "chkpt", "quazimoto_uld.pt"))
69
+ p.add_argument("--ckpt-every", type=int, default=250)
70
+ p.add_argument("--resume", action="store_true")
71
+ # streaming blend (prose source) -- same flags/defaults as train.py
72
+ p.add_argument("--stream", action="store_true")
73
+ p.add_argument("--data", default="")
74
+ for name, d in [("ultra-frac", 0.35), ("edu-frac", 0.25),
75
+ ("math-frac", 0.25), ("pretrain-frac", 0.15)]:
76
+ p.add_argument(f"--{name}", type=float, default=d)
77
+ p.add_argument("--ultra-dataset", default="openbmb/Ultra-FineWeb-L3")
78
+ p.add_argument("--ultra-config", default="Ultra-FineWeb-L3-en-Multi-Style-Synthetic")
79
+ p.add_argument("--fineweb-dataset", default="HuggingFaceFW/fineweb-edu")
80
+ p.add_argument("--fineweb-config", default="sample-10BT")
81
+ p.add_argument("--math-dataset", default="HuggingFaceTB/finemath")
82
+ p.add_argument("--math-config", default="finemath-4plus")
83
+ p.add_argument("--pretrain-dataset", default="Quazim0t0/PretrainNew")
84
+ p.add_argument("--pretrain-config", default=None)
85
+ p.add_argument("--split", default="train")
86
+ p.add_argument("--seed", type=int, default=0)
87
+ args = p.parse_args()
88
+
89
+ torch.manual_seed(args.seed)
90
+ dev, tdev = args.device, (args.teacher_device or args.device)
91
+ s_tok = T.load_tokenizer(args.tok_dir)
92
+ s_pad = s_tok.pad_token_id or 0
93
+
94
+ ck = torch.load(args.student_ckpt, map_location=dev, weights_only=False)
95
+ cfg = QuazimotoConfig(**ck["family_config"])
96
+ V = cfg.vocab_size
97
+ student = QuazimotoLM(cfg); student.load_state_dict(ck["model"], strict=False)
98
+ student.to(dev).train()
99
+ print(f"student: {args.student_ckpt} (step {ck.get('step')}), vocab {V}")
100
+
101
+ from transformers import AutoModelForCausalLM, AutoTokenizer
102
+ t_tok = AutoTokenizer.from_pretrained(args.teacher, trust_remote_code=True)
103
+ try:
104
+ teacher = AutoModelForCausalLM.from_pretrained(args.teacher, dtype=torch.bfloat16,
105
+ trust_remote_code=True)
106
+ except TypeError:
107
+ teacher = AutoModelForCausalLM.from_pretrained(args.teacher, torch_dtype=torch.bfloat16,
108
+ trust_remote_code=True)
109
+ teacher = teacher.to(tdev).eval()
110
+ for pp in teacher.parameters():
111
+ pp.requires_grad_(False)
112
+ t_pad = t_tok.pad_token_id if t_tok.pad_token_id is not None else t_tok.eos_token_id
113
+ print(f"teacher: {args.teacher} on {tdev} (frozen, bf16)")
114
+
115
+ # ---- data: same text tokenized two ways ----
116
+ if args.stream:
117
+ docs = T.doc_generator(args)
118
+ get_text = lambda: next(docs)
119
+ else:
120
+ ids_all, _ = T.load_ids(args.data, s_tok)
121
+ def get_text():
122
+ import numpy as np
123
+ i = np.random.randint(0, max(len(ids_all) - args.seq, 1))
124
+ return s_tok.decode(ids_all[i:i + args.seq].tolist(), skip_special_tokens=True)
125
+
126
+ def get_batch():
127
+ texts = [get_text() for _ in range(args.micro)]
128
+ s_rows = [s_tok.encode(t, add_special_tokens=False)[:args.seq] for t in texts]
129
+ t_rows = [t_tok(t, add_special_tokens=False).input_ids[:args.seq] for t in texts]
130
+ sL = max(len(r) for r in s_rows); tL = max(len(r) for r in t_rows)
131
+ s_ids = torch.full((args.micro, sL), s_pad, dtype=torch.long)
132
+ targets = torch.full((args.micro, sL), -1, dtype=torch.long) # ignore_index=-1
133
+ t_ids = torch.full((args.micro, tL), t_pad, dtype=torch.long)
134
+ uld_n = []
135
+ for j, (sr, tr) in enumerate(zip(s_rows, t_rows)):
136
+ s_ids[j, :len(sr)] = torch.tensor(sr)
137
+ if len(sr) > 1:
138
+ targets[j, :len(sr) - 1] = torch.tensor(sr[1:]) # next-token targets
139
+ t_ids[j, :len(tr)] = torch.tensor(tr)
140
+ uld_n.append(min(len(sr), len(tr)) - 1)
141
+ return (s_ids.to(dev), targets.to(dev), t_ids.to(tdev),
142
+ (t_ids != t_pad).to(tdev), uld_n)
143
+
144
+ opt = torch.optim.AdamW(student.parameters(), lr=args.lr, betas=(0.9, 0.95),
145
+ weight_decay=args.weight_decay)
146
+
147
+ start_step = 1
148
+ if args.resume and os.path.isfile(args.out):
149
+ r = torch.load(args.out, map_location=dev, weights_only=False)
150
+ student.load_state_dict(r["model"], strict=False)
151
+ if "optim" in r:
152
+ try: opt.load_state_dict(r["optim"])
153
+ except ValueError: pass
154
+ start_step = int(r.get("step", 0)) + 1
155
+ print(f"resumed ULD from {args.out} at step {r.get('step')} -> {start_step}")
156
+ if start_step > args.steps:
157
+ print(f" nothing to do: already at {start_step-1} >= --steps {args.steps}."); return
158
+
159
+ t0 = time.time()
160
+ for step in range(start_step, args.steps + 1):
161
+ for g in opt.param_groups:
162
+ g["lr"] = T.lr_at(step, args.lr, args.warmup, args.steps, args.min_lr_frac)
163
+ opt.zero_grad(set_to_none=True)
164
+ tot = tul = tce = 0.0
165
+ for _ in range(args.accum):
166
+ s_ids, targets, t_ids, t_mask, uld_n = get_batch()
167
+ with torch.no_grad():
168
+ t_logits = teacher(t_ids, attention_mask=t_mask).logits.to(dev)
169
+ with torch.autocast("cuda", torch.bfloat16, enabled=args.amp and dev == "cuda"):
170
+ logits, _, _ = student(s_ids) # targets=None -> no aux
171
+ ce = F.cross_entropy(logits.reshape(-1, V), targets.reshape(-1), ignore_index=-1)
172
+ ul = uld_loss(logits, t_logits, uld_n, args.top_k)
173
+ loss = ce + args.uld_w * ul
174
+ (loss / args.accum).backward()
175
+ tot += loss.item() / args.accum; tul += ul.item() / args.accum; tce += ce.item() / args.accum
176
+ del t_logits, logits
177
+ gn = torch.nn.utils.clip_grad_norm_(student.parameters(), 1.0)
178
+ if torch.isfinite(gn):
179
+ opt.step()
180
+ else:
181
+ print(f"step {step}: non-finite grad, skip")
182
+
183
+ if step % 25 == 0 or step == 1:
184
+ dt = (time.time() - t0) / (step - start_step + 1)
185
+ print(f"step {step}/{args.steps} | loss {tot:.3f} | ce {tce:.3f} | uld {tul:.3f} "
186
+ f"| ppl {math.exp(min(tce,20)):.1f} | {dt:.1f}s/step", flush=True)
187
+ if args.ckpt_every and step % args.ckpt_every == 0:
188
+ T.save_ckpt(student, s_tok.vocab_size, step, args.out, opt)
189
+ T.save_ckpt(student, s_tok.vocab_size, args.steps, args.out, opt)
190
+ print("done ->", args.out)
191
+
192
+
193
+ if __name__ == "__main__":
194
+ main()
family.py ADDED
@@ -0,0 +1,446 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """SpikeWhale / Byrne family traits, ported to Quazimoto-LM (attention backbone).
2
+
3
+ These are the *transformer-native* family blocks (their origin is the transformer
4
+ `modeling_byrne_embed.py`), so unlike the SNN port they operate directly on the
5
+ sequence hidden state [B,T,d] -- no per-step adaptation needed. Each keeps the
6
+ family's safe-at-init contract: a tanh/zero gate makes the block a no-op at start,
7
+ while the content (`up`/`down`) weights are NON-zero so the gate still receives
8
+ gradient (the double-zero saddle would deadlock it). DERF soft_clamp bounds any
9
+ new instability surface, in line with the family's stability discipline.
10
+
11
+ Included: HRMRefinementBlock (signature), MoESwiGLU, MTPHead, JEPAPredictorBlock.
12
+ Engram / ProgSem are bio/SNN-specific and SpikingLinearAttention is the SNN's
13
+ stand-in for the real attention this model already has, so they are omitted.
14
+ """
15
+ from __future__ import annotations
16
+
17
+ import math
18
+ import torch
19
+ import torch.nn as nn
20
+ import torch.nn.functional as F
21
+
22
+ import instrument as _viz # live-visualizer capture hooks (no-op unless a recorder is active)
23
+
24
+ # sqrt(pi)/2: soft_clamp is the identity for small inputs and saturates smoothly to
25
+ # +/-bound with a non-zero gradient everywhere (no dead-gradient zones).
26
+ _ERF_K = math.sqrt(math.pi) / 2.0
27
+
28
+
29
+ def soft_clamp(x, bound):
30
+ return bound * torch.erf(x * (_ERF_K / bound))
31
+
32
+
33
+ class RMSNorm(nn.Module):
34
+ def __init__(self, dim, eps=1e-6):
35
+ super().__init__()
36
+ self.eps = eps
37
+ self.weight = nn.Parameter(torch.ones(dim))
38
+
39
+ def forward(self, x):
40
+ return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) * self.weight
41
+
42
+
43
+ def sqrtsoftplus(x):
44
+ """Family expert-scoring function: sqrt(softplus(x))."""
45
+ return torch.sqrt(F.softplus(x) + 1e-8)
46
+
47
+
48
+ class HRMRefinementBlock(nn.Module):
49
+ """Signature family block: iterative gated refinement that, in the canonical
50
+ HRM spirit, starts the reasoning from a RANDOM initial state z0 and iterates
51
+ toward a solution conditioned on the input (`anchor`).
52
+
53
+ Unlike the other family traits this block does NOT start as a no-op: the gate
54
+ is initialised OPEN (gate_init_open). With an input-anchored no-op start the
55
+ gate received ~zero gradient and never woke up; a random z0 forces the block
56
+ to actively reconcile the random state against the input, so the open gate
57
+ carries real signal from step 0. To keep the deep trunk intact we contribute
58
+ only the reasoning DELTA (h - z0) as a residual -- z0 itself is never dumped
59
+ into the trunk, and a closed gate (h == z0) degrades cleanly to a no-op."""
60
+
61
+ def __init__(self, hidden_size, refine_dim, steps, eps=1e-3, gate_init_open=0.1):
62
+ super().__init__()
63
+ self.steps = steps
64
+ self.norm = RMSNorm(hidden_size, eps)
65
+ self.down = nn.Linear(hidden_size * 2, refine_dim, bias=False)
66
+ self.up = nn.Linear(refine_dim, hidden_size, bias=False)
67
+ # random initial reasoning state (learnable), broadcast over batch/time
68
+ self.z0 = nn.Parameter(torch.empty(hidden_size))
69
+ nn.init.trunc_normal_(self.z0, std=1.0, a=-2.0, b=2.0)
70
+ # gates start OPEN so the random-state reasoning reaches the output at init
71
+ go = math.atanh(min(gate_init_open, 0.9)) if gate_init_open > 0 else 0.0
72
+ self.gate = nn.Parameter(torch.full((steps,), go))
73
+ nn.init.normal_(self.down.weight, std=0.02)
74
+ nn.init.normal_(self.up.weight, std=0.02)
75
+
76
+ def forward(self, x): # x: [B,T,d]
77
+ B, T, _ = x.shape
78
+ anchor = x
79
+ h = self.z0.expand(B, T, -1) # random initial reasoning state
80
+ for t in range(self.steps):
81
+ inp = torch.cat([self.norm(h), anchor], dim=-1)
82
+ update = soft_clamp(self.up(F.silu(self.down(inp))), 10.0)
83
+ h = h + torch.tanh(self.gate[t]) * update
84
+ return x + (h - self.z0) # add reasoning delta, keep trunk
85
+
86
+
87
+ class ExpertFFN(nn.Module):
88
+ def __init__(self, hidden_size, intermediate_size):
89
+ super().__init__()
90
+ self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
91
+ self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
92
+ self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False)
93
+
94
+ def forward(self, x):
95
+ return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
96
+
97
+
98
+ class MoESwiGLU(nn.Module):
99
+ """Shared + top-k routed SwiGLU experts, sqrtsoftplus scoring, norm_topk_prob,
100
+ Switch-style load-balance aux (via `last_aux_loss`). down-projections zero-init
101
+ => no-op at start."""
102
+
103
+ def __init__(self, hidden_size, intermediate_size, n_routed=4, n_shared=1,
104
+ top_k=2, aux_loss_coef=0.01):
105
+ super().__init__()
106
+ self.top_k = min(top_k, n_routed)
107
+ self.n_routed = n_routed
108
+ self.n_shared = n_shared
109
+ self.aux_loss_coef = aux_loss_coef
110
+ self.router = nn.Linear(hidden_size, n_routed, bias=False)
111
+ self.experts = nn.ModuleList([ExpertFFN(hidden_size, intermediate_size)
112
+ for _ in range(n_routed)])
113
+ self.shared = (ExpertFFN(hidden_size, intermediate_size * n_shared)
114
+ if n_shared > 0 else None)
115
+ for e in self.experts:
116
+ nn.init.zeros_(e.down_proj.weight)
117
+ if self.shared is not None:
118
+ nn.init.zeros_(self.shared.down_proj.weight)
119
+ self.last_aux_loss = None
120
+
121
+ def forward(self, x): # x: [B,T,d]
122
+ flat = x.reshape(-1, x.shape[-1])
123
+ out = torch.zeros_like(flat)
124
+ if self.shared is not None:
125
+ s = self.shared(flat)
126
+ out = out + (s / self.n_shared if self.n_shared > 1 else s)
127
+
128
+ logits = self.router(flat)
129
+ scores = sqrtsoftplus(logits)
130
+ topv, topi = scores.topk(self.top_k, dim=-1)
131
+ topv = topv / (topv.sum(-1, keepdim=True) + 1e-8)
132
+ for slot in range(self.top_k):
133
+ idx = topi[:, slot]
134
+ w = topv[:, slot].unsqueeze(-1)
135
+ for e_id, expert in enumerate(self.experts):
136
+ mask = idx == e_id
137
+ if mask.any():
138
+ out[mask] = out[mask] + w[mask] * expert(flat[mask])
139
+
140
+ probs = F.softmax(logits, dim=-1)
141
+ expert_mask = torch.zeros_like(probs)
142
+ expert_mask.scatter_(1, topi, 1.0)
143
+ self.last_aux_loss = (self.n_routed * (expert_mask.mean(0) * probs.mean(0)).sum()
144
+ * self.aux_loss_coef)
145
+ return out.view_as(x)
146
+
147
+
148
+ class MTPHead(nn.Module):
149
+ """Multi-token-prediction head: zero-init d->d residual, reuses the tied readout."""
150
+
151
+ def __init__(self, hidden_size):
152
+ super().__init__()
153
+ self.proj = nn.Linear(hidden_size, hidden_size, bias=False)
154
+ nn.init.zeros_(self.proj.weight)
155
+
156
+ def forward(self, hidden):
157
+ return hidden + self.proj(hidden)
158
+
159
+
160
+ class TokenCompressor(nn.Module):
161
+ """Frozen LSH-style projection (gradient never reaches it through the hash cast)."""
162
+ def __init__(self, hidden_size, compress_dim):
163
+ super().__init__()
164
+ self.proj = nn.Linear(hidden_size, compress_dim, bias=False)
165
+ nn.init.normal_(self.proj.weight, std=0.02)
166
+ self.proj.weight.requires_grad_(False)
167
+
168
+ def forward(self, x):
169
+ return self.proj(x)
170
+
171
+
172
+ class MultiHeadHashLookup(nn.Module):
173
+ """N-gram hash memory: for n=1..max_ngram, hash the n-token compressed window
174
+ into per-head tables and average. Ported from v2 EngramModule."""
175
+ def __init__(self, num_heads, table_size, compress_dim, out_dim, max_ngram=3):
176
+ super().__init__()
177
+ self.num_heads, self.table_size = num_heads, table_size
178
+ self.max_ngram, self.out_dim = max_ngram, out_dim
179
+ self.tables = nn.ModuleList([nn.Embedding(table_size, out_dim) for _ in range(num_heads)])
180
+ for t in self.tables:
181
+ nn.init.normal_(t.weight, std=0.01)
182
+ for n in range(1, max_ngram + 1):
183
+ for k in range(n):
184
+ proj = torch.randn(num_heads, compress_dim)
185
+ proj = proj / (proj.norm(dim=1, keepdim=True) + 1e-8)
186
+ self.register_buffer(f"hash_proj_n{n}_p{k}", proj, persistent=True)
187
+
188
+ def forward(self, compressed):
189
+ B, S, _ = compressed.shape
190
+ dev = compressed.device
191
+ out = torch.zeros(B, S, self.out_dim, device=dev, dtype=compressed.dtype)
192
+ norm = torch.zeros(S, device=dev)
193
+ for n in range(1, self.max_ngram + 1):
194
+ if S < n:
195
+ continue
196
+ valid, start = S - n + 1, n - 1
197
+ h = torch.zeros(B, valid, self.num_heads, device=dev)
198
+ for k in range(n):
199
+ proj = getattr(self, f"hash_proj_n{n}_p{k}")
200
+ h = h + torch.matmul(compressed[:, k:k + valid, :].float(), proj.t())
201
+ idx = h.abs().long() % self.table_size
202
+ for hi, table in enumerate(self.tables):
203
+ out[:, start:, :] = out[:, start:, :] + table(idx[:, :, hi])
204
+ norm[start:] += self.num_heads
205
+ return (out / norm.view(1, -1, 1).clamp(min=1)).to(compressed.dtype)
206
+
207
+
208
+ class DERFContextGate(nn.Module):
209
+ def __init__(self, obs_size, init_bias=-4.0):
210
+ super().__init__()
211
+ self.proj = nn.Linear(obs_size * 2, obs_size)
212
+ self.alpha = nn.Parameter(torch.ones(obs_size))
213
+ self.bias = nn.Parameter(torch.full((obs_size,), init_bias))
214
+ self.gamma = nn.Parameter(torch.ones(obs_size))
215
+
216
+ def forward(self, retrieved, obs):
217
+ logits = self.proj(torch.cat([retrieved, obs], dim=-1))
218
+ gate = self.gamma * ((torch.erf(self.alpha * logits + self.bias) + 1.0) / 2.0)
219
+ return retrieved * gate
220
+
221
+
222
+ class PhaseAttentionRing(nn.Module):
223
+ """Interstitial ATTENTION ring: attends causally over the sequence in
224
+ oscillator-PHASE space ([cos,sin] of the two neighbor oscillator rings) and
225
+ returns an injection current of width m = n_r + n_{r+1} for those neighbors.
226
+ Zero-init gate => no-op at start; soft_clamp bounds the injected drive."""
227
+
228
+ def __init__(self, m, n_heads=4, head_dim=16, bound=10.0):
229
+ super().__init__()
230
+ self.h, self.d, self.bound = n_heads, head_dim, bound
231
+ self.qkv = nn.Linear(2 * m, 3 * n_heads * head_dim, bias=False)
232
+ self.out = nn.Linear(n_heads * head_dim, m, bias=False)
233
+ self.gate = nn.Parameter(torch.zeros(1))
234
+ nn.init.normal_(self.qkv.weight, std=0.02)
235
+ nn.init.normal_(self.out.weight, std=0.02)
236
+
237
+ def forward(self, theta_slice): # [B,T,m] phases of the neighbor rings
238
+ B, T, m = theta_slice.shape
239
+ feat = torch.cat([torch.cos(theta_slice), torch.sin(theta_slice)], dim=-1)
240
+ q, k, v = self.qkv(feat).split(self.h * self.d, dim=-1)
241
+ shp = lambda z: z.view(B, T, self.h, self.d).transpose(1, 2)
242
+ y = F.scaled_dot_product_attention(shp(q), shp(k), shp(v), is_causal=True)
243
+ y = y.transpose(1, 2).reshape(B, T, self.h * self.d)
244
+ return soft_clamp(self.out(y) * torch.tanh(self.gate), self.bound)
245
+
246
+
247
+ class EngramRing(nn.Module):
248
+ """Interstitial ENGRAM ring: absorbs n-gram context from the hidden state via
249
+ hash memory + DERF gate, projected to an injection current of width m for the
250
+ two neighbor oscillator rings. Two no-op gates at init (DERF bias -4 + scale)."""
251
+
252
+ def __init__(self, hidden_size, m, compress_dim=32, num_heads=2,
253
+ table_size=2048, max_ngram=3, bound=10.0):
254
+ super().__init__()
255
+ self.bound = bound
256
+ self.compressor = TokenCompressor(hidden_size, compress_dim)
257
+ self.lookup = MultiHeadHashLookup(num_heads, table_size, compress_dim, m, max_ngram)
258
+ self.to_obs = nn.Linear(hidden_size, m, bias=False)
259
+ self.gate = DERFContextGate(m, init_bias=-4.0)
260
+ self.scale = nn.Parameter(torch.zeros(1)) # extra no-op gate at init
261
+ nn.init.normal_(self.to_obs.weight, std=0.02)
262
+
263
+ def family_reinit(self):
264
+ """Re-apply the inits the model's global self.apply would clobber (table
265
+ std 0.01, frozen-random compressor, DERF bias -4)."""
266
+ for t in self.lookup.tables:
267
+ nn.init.normal_(t.weight, std=0.01)
268
+ nn.init.normal_(self.compressor.proj.weight, std=0.02)
269
+ self.compressor.proj.weight.requires_grad_(False)
270
+ nn.init.constant_(self.gate.bias, -4.0)
271
+
272
+ def forward(self, h): # h: [B,T,hidden]
273
+ retrieved = self.lookup(self.compressor(h.detach()))
274
+ gated = self.gate(retrieved, self.to_obs(h))
275
+ return soft_clamp(gated * torch.tanh(self.scale), self.bound)
276
+
277
+
278
+ class RingController(nn.Module):
279
+ """Tiny per-ring manager that OPTIMIZES ITSELF by a predictive / free-energy rule.
280
+
281
+ Core = a fast-weight linear predictor `W` (a BUFFER, excluded from the global
282
+ optimizer) updated online by the delta rule W += lr * (f - W@prev) outer prev,
283
+ which is exactly one gradient step on the squared prediction error -- the
284
+ controller learns to predict its ring's next state, minimizing surprise, with no
285
+ backprop. A small backprop-trained decoder maps the self-organized feature +
286
+ surprise into ring-control modulations; zero-init => exact no-op at start."""
287
+
288
+ def __init__(self, d_obs=4, feat=384, n_ctrl=4, local_lr=0.01):
289
+ super().__init__()
290
+ self.feat, self.local_lr = feat, local_lr
291
+ self.enc = nn.Linear(d_obs, feat)
292
+ self.dec = nn.Linear(feat + 1, n_ctrl)
293
+ nn.init.zeros_(self.dec.weight)
294
+ nn.init.zeros_(self.dec.bias) # control == 0 at init (no-op)
295
+ self.register_buffer("W", torch.zeros(feat, feat)) # self-organizing fast weights
296
+ self.register_buffer("prev_f", torch.zeros(feat))
297
+
298
+ def family_reinit(self):
299
+ nn.init.zeros_(self.dec.weight)
300
+ nn.init.zeros_(self.dec.bias)
301
+
302
+ def forward(self, obs): # obs: [d_obs] (detached ring stats)
303
+ f = torch.tanh(self.enc(obs)) # [feat]
304
+ pred = self.W @ self.prev_f # predicted current feature
305
+ surprise = F.mse_loss(f.detach(), pred)
306
+ if self.training:
307
+ with torch.no_grad(): # predictive self-organization (no global grad)
308
+ err = f.detach() - pred
309
+ self.W.add_(self.local_lr * torch.outer(err, self.prev_f)).clamp_(-3.0, 3.0)
310
+ self.prev_f.copy_(f.detach())
311
+ ctrl = self.dec(torch.cat([f, surprise.detach().reshape(1)])) # [n_ctrl]
312
+ return ctrl, surprise.detach()
313
+
314
+
315
+ class RingControllerBank(nn.Module):
316
+ """One RingController per oscillator ring (shared across all layers)."""
317
+
318
+ def __init__(self, n_rings, d_obs=4, feat=384, local_lr=0.01):
319
+ super().__init__()
320
+ self.controllers = nn.ModuleList(
321
+ [RingController(d_obs, feat, 4, local_lr) for _ in range(n_rings)])
322
+ self.last_surprise = None
323
+
324
+ def forward(self, obs): # obs: [R, d_obs] -> ctrl [R, 4]
325
+ ctrls, surps = [], []
326
+ for r, c in enumerate(self.controllers):
327
+ ct, sp = c(obs[r])
328
+ ctrls.append(ct)
329
+ surps.append(sp)
330
+ self.last_surprise = torch.stack(surps).mean()
331
+ return torch.stack(ctrls, dim=0)
332
+
333
+
334
+ class RingSpecialists(nn.Module):
335
+ """A MoE-style bank of `n_spec` MINI MEMORY SPECIALISTS for ONE oscillator ring.
336
+
337
+ Each specialist owns two fast-weight stores (test-time-mutable BUFFERS, like
338
+ RingController.W -- excluded from the optimizer):
339
+ * store_in -- a memory of the INPUT context that routes to it, and
340
+ * store_out -- the OUTPUT information it injects back into the ring.
341
+ Tokens are routed to the top-k specialists (a small MoE) by similarity to each
342
+ specialist's address = its learnable identity key + a read of its accumulated
343
+ input memory. The routed store_out is decoded into an injection current for the
344
+ ring, and BOTH stores are written online (train AND inference) by a gated EMA
345
+ rule -- so a generation accumulates an addressable context memory as it runs.
346
+
347
+ Slow (backprop) weights -- q_proj/in_enc/val_enc/out_dec/in_read/key/active --
348
+ learn to route, encode, retrieve and decode; the stores are the fast memory.
349
+ Family contract: zero-init `scale` => exact no-op at start, and an empty
350
+ store_out is zero anyway, so the block is doubly safe until it learns to write
351
+ and open the gate. `active` is a per-specialist usage gate biasing the router."""
352
+
353
+ def __init__(self, ring_size, hidden_size, n_spec=7, key_dim=32, slot_dim=64,
354
+ top_k=2, write_lr=0.1, bound=10.0):
355
+ super().__init__()
356
+ self.n_spec = n_spec
357
+ self.top_k = min(top_k, n_spec)
358
+ self.write_lr, self.bound = write_lr, bound
359
+ self.write_enabled = True
360
+ self.q_proj = nn.Linear(hidden_size, key_dim, bias=False) # router query (learned)
361
+ self.out_dec = nn.Linear(slot_dim, ring_size, bias=False) # retrieved -> injection (learned)
362
+ self.in_read = nn.Linear(slot_dim, key_dim, bias=False) # input-store -> addr (learned)
363
+ # write-side encoders are FROZEN RANDOM projections (cf. EngramRing's frozen
364
+ # compressor): they only ever run inside the no-grad write, so backprop can't
365
+ # train them -- as fixed random features the stores hold a stable encoding the
366
+ # learned read/route/decode path can address.
367
+ self.in_enc = nn.Linear(hidden_size, slot_dim, bias=False) # input -> input-store (frozen)
368
+ self.val_enc = nn.Linear(hidden_size, slot_dim, bias=False) # input -> output-store (frozen)
369
+ self.key = nn.Parameter(torch.randn(n_spec, key_dim) * 0.02) # specialist identity
370
+ self.active = nn.Parameter(torch.zeros(n_spec)) # per-specialist usage gate
371
+ self.scale = nn.Parameter(torch.zeros(1)) # no-op output gate at init
372
+ self.register_buffer("store_in", torch.zeros(n_spec, slot_dim))
373
+ self.register_buffer("store_out", torch.zeros(n_spec, slot_dim))
374
+ for m in (self.q_proj, self.out_dec, self.in_read, self.in_enc, self.val_enc):
375
+ nn.init.normal_(m.weight, std=0.02)
376
+ self.in_enc.weight.requires_grad_(False)
377
+ self.val_enc.weight.requires_grad_(False)
378
+
379
+ def family_reinit(self):
380
+ """Re-apply inits the model's global self.apply clobbers (key, gates, and
381
+ the frozen write encoders)."""
382
+ nn.init.normal_(self.key, std=0.02)
383
+ nn.init.zeros_(self.active)
384
+ nn.init.zeros_(self.scale)
385
+ nn.init.normal_(self.in_enc.weight, std=0.02)
386
+ nn.init.normal_(self.val_enc.weight, std=0.02)
387
+ self.in_enc.weight.requires_grad_(False)
388
+ self.val_enc.weight.requires_grad_(False)
389
+
390
+ def reset_memory(self):
391
+ """Clear both stores -- call between independent prompts/sequences so
392
+ context memory does not bleed across them."""
393
+ self.store_in.zero_()
394
+ self.store_out.zero_()
395
+
396
+ def forward(self, h): # h: [B,T,hidden]
397
+ B, T, _ = h.shape
398
+ # snapshot the fast-weight stores: the graph must hold an immutable copy
399
+ # because we mutate the buffers in-place for the online write below.
400
+ store_in, store_out = self.store_in.clone(), self.store_out.clone()
401
+ q = self.q_proj(h) # [B,T,key_dim]
402
+ addr = self.key + self.in_read(store_in) # [n_spec,key_dim]
403
+ logits = q @ addr.t() # [B,T,n_spec]
404
+ logits = logits + F.logsigmoid(self.active) # usage gate biases routing
405
+ if self.top_k < self.n_spec: # top-k MoE sparsity
406
+ tv = torch.topk(logits, self.top_k, dim=-1).values
407
+ logits = logits.masked_fill(logits < tv[..., [-1]], float("-inf"))
408
+ route = torch.softmax(logits, dim=-1) # [B,T,n_spec]
409
+
410
+ retrieved = route @ store_out # [B,T,slot_dim]
411
+ inject = self.out_dec(retrieved) * torch.tanh(self.scale) # [B,T,ring_size]
412
+
413
+ rec = _viz.get_rec()
414
+ if rec is not None and rec.enabled: # last-token routing
415
+ rec.push_spec(route[0, -1].tolist())
416
+
417
+ # online write: blend this step's input/value into the routed specialists
418
+ if self.write_enabled and self.write_lr > 0:
419
+ with torch.no_grad():
420
+ w = route.reshape(-1, self.n_spec) # [BT,n_spec]
421
+ denom = w.sum(0).clamp(min=1e-3).unsqueeze(1) # [n_spec,1]
422
+ in_info = (w.t() @ self.in_enc(h).reshape(-1, self.in_enc.out_features)) / denom
423
+ val_info = (w.t() @ self.val_enc(h).reshape(-1, self.val_enc.out_features)) / denom
424
+ a = self.write_lr
425
+ self.store_in.mul_(1 - a).add_(a * in_info).clamp_(-self.bound, self.bound)
426
+ self.store_out.mul_(1 - a).add_(a * val_info).clamp_(-self.bound, self.bound)
427
+ return soft_clamp(inject, self.bound)
428
+
429
+
430
+ class JEPAPredictorBlock(nn.Module):
431
+ """Representation-space k-ahead prediction with stop-grad target (JEPA asymmetry).
432
+ Zero-init gate => identity at init; `up` normal so the gate gets gradient."""
433
+
434
+ def __init__(self, dim, pred_dim, horizon, eps=1e-3):
435
+ super().__init__()
436
+ self.horizon = horizon
437
+ self.norm = RMSNorm(dim, eps)
438
+ self.down = nn.Linear(dim, pred_dim, bias=False)
439
+ self.up = nn.Linear(pred_dim, dim, bias=False)
440
+ self.gate = nn.Parameter(torch.zeros(horizon))
441
+ nn.init.normal_(self.down.weight, std=0.02)
442
+ nn.init.normal_(self.up.weight, std=0.02)
443
+
444
+ def forward(self, h, k): # h: [B,T,dim]
445
+ update = self.up(F.silu(self.down(self.norm(h))))
446
+ return h + torch.tanh(self.gate[k - 1]) * update
fractal.py ADDED
@@ -0,0 +1,79 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ fractal.py -- Mandelbrot phase seeding for Quazimoto-LM's oscillator bank.
3
+
4
+ The recommended (and only) fractal integration: instead of a generic learned
5
+ phase initializer, give each TOKEN a characteristic dynamical signature drawn
6
+ from the Mandelbrot iteration z <- z^2 + c, and use the ANGLES of that orbit to
7
+ seed the N oscillator phases. The Mandelbrot map is itself an iterated dynamical
8
+ system, so this hands the Kuramoto ring bank a token-specific, deterministic,
9
+ parameter-free phase prior congruent with what the block already does.
10
+
11
+ We build a frozen [vocab_size, n_osc] table once: token id -> a complex point c
12
+ (spread over the Mandelbrot region by a 2D Halton low-discrepancy sequence so
13
+ coverage is even and deterministic) -> phase_k = angle(z_k) for the first n_osc
14
+ orbit points. The model adds this, through a zero-init gate, to to_theta(h) inside
15
+ each QuazimotoBlock -- a no-op at init that the optimizer can choose to open.
16
+
17
+ Smooth-by-construction: we read the orbit ANGLE (always defined, bounded to
18
+ (-pi, pi]) rather than escape-time, avoiding the chaotic boundary discontinuities
19
+ that raw escape counts would inject.
20
+ """
21
+
22
+ import torch
23
+
24
+
25
+ def _halton(i, base):
26
+ """Radical-inverse (van der Corput) value of i in the given base, in [0,1)."""
27
+ f, r = 1.0, 0.0
28
+ while i > 0:
29
+ f /= base
30
+ r += f * (i % base)
31
+ i //= base
32
+ return r
33
+
34
+
35
+ @torch.no_grad()
36
+ def mandelbrot_phase_table(vocab_size, n_osc, region=(-2.5, 1.0, -1.25, 1.25),
37
+ clamp_mag=1e3):
38
+ """Return a frozen [vocab_size, n_osc] tensor of orbit-angle phase seeds.
39
+
40
+ token id -> c via 2D Halton(base 2,3) over `region`; phase_k = angle(z_k) for
41
+ k = 0..n_osc-1 of the iteration z <- z^2 + c (z0 = 0). This is the FLAT
42
+ (tokenizer-agnostic) map; for the hierarchical byte-merge map build the table
43
+ offline with build_fractal_table.py and load it via load_phase_table()."""
44
+ x0, x1, y0, y1 = region
45
+ ids = torch.arange(1, vocab_size + 1)
46
+ hx = torch.tensor([_halton(int(i), 2) for i in ids])
47
+ hy = torch.tensor([_halton(int(i), 3) for i in ids])
48
+ cr = x0 + hx * (x1 - x0)
49
+ ci = y0 + hy * (y1 - y0)
50
+ return phases_from_c(torch.complex(cr, ci), n_osc, clamp_mag)
51
+
52
+
53
+ @torch.no_grad()
54
+ def phases_from_c(c, n_osc, clamp_mag=1e3):
55
+ """Orbit-angle phases for a batch of complex seeds c [V] -> [V, n_osc].
56
+ phase_k = angle(z_k) of z <- z^2 + c (z0=0); magnitude clamped (angle kept)."""
57
+ z = torch.zeros_like(c)
58
+ phases = torch.empty(c.shape[0], n_osc)
59
+ for k in range(n_osc):
60
+ z = z * z + c
61
+ mag = z.abs()
62
+ over = mag > clamp_mag # rescale escaped orbits, keep direction
63
+ if over.any():
64
+ z = torch.where(over, z / mag * clamp_mag, z)
65
+ phases[:, k] = torch.angle(z) # in (-pi, pi]
66
+ return phases
67
+
68
+
69
+ def load_phase_table(vocab_size, n_osc, path):
70
+ """Load a precomputed phase table if it matches (vocab_size, n_osc); else None.
71
+ Returns (phases, mode) or (None, None)."""
72
+ import os
73
+ if not os.path.exists(path):
74
+ return None, None
75
+ d = torch.load(path, map_location="cpu", weights_only=False)
76
+ ph = d.get("phases")
77
+ if ph is not None and tuple(ph.shape) == (vocab_size, n_osc):
78
+ return ph.float(), d.get("mode", "precomputed")
79
+ return None, None
fractal_phase.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:ee042ce925f1a47e0a7b129fb78f95ef2f5fe979cfd2bb9e3b410185e8dea952
3
+ size 19023162
generate.py ADDED
@@ -0,0 +1,326 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ generate.py -- inference / generation harness for Quazimoto-LM.
3
+
4
+ Loads a checkpoint saved by train.py (save_ckpt writes
5
+ {"model", "family_config", "vocab_size", "step"})
6
+ rebuilds the exact model from the embedded family_config, restores weights, and
7
+ samples completions with the bundled SpikeWhale tokenizer.
8
+
9
+ This is tuned to the Quazimoto architecture's quirks:
10
+ * The model has NO KV cache (attention is training-only / full recompute), so
11
+ each new token re-runs the whole context. We respect cfg.block_size by
12
+ cropping the context window, exactly like model.generate().
13
+ * We reuse the model's NaN/inf-sanitised sampling path (the oscillator stack
14
+ can push logits to +/-inf over long rollouts), and add nucleus (top-p),
15
+ top-k, and repetition-penalty controls on top.
16
+ * Generation stops early on <eos> / <|im_end|> / <|endoftext|> when present.
17
+
18
+ Examples
19
+ --------
20
+ # plain completion
21
+ python generate.py --ckpt chkpt/quazimoto.pt --prompt "the quazimoto oscillator"
22
+
23
+ # chat-style turn (wraps prompt in ChatML and stops on <|im_end|>)
24
+ python generate.py --ckpt chkpt/quazimoto.pt --chat --prompt "Hello, who are you?"
25
+
26
+ # interactive REPL
27
+ python generate.py --ckpt chkpt/quazimoto.pt --interactive
28
+ """
29
+
30
+ import argparse
31
+ import os
32
+ import sys
33
+
34
+ import torch
35
+ import torch.nn.functional as F
36
+
37
+ from model import QuazimotoLM, QuazimotoConfig
38
+
39
+ PKG_DIR = os.path.dirname(os.path.abspath(__file__))
40
+
41
+
42
+ # --------------------------------------------------------------------------- #
43
+ # loading
44
+ # --------------------------------------------------------------------------- #
45
+ def load_tokenizer(tok_dir):
46
+ sys.path.insert(0, tok_dir)
47
+ from spike_tokenizer import SpikeTokenizer
48
+ return SpikeTokenizer(vocab_file=os.path.join(tok_dir, "tokenizer.json"))
49
+
50
+
51
+ def load_model(ckpt_path, device):
52
+ ckpt = torch.load(ckpt_path, map_location=device, weights_only=False)
53
+ fam = ckpt.get("family_config")
54
+ if fam is None:
55
+ raise ValueError(f"{ckpt_path} has no family_config; can't rebuild the model.")
56
+ cfg = QuazimotoConfig(**fam)
57
+ model = QuazimotoLM(cfg)
58
+ state = ckpt["model"]
59
+ missing, unexpected = model.load_state_dict(state, strict=False)
60
+ if missing:
61
+ print(f" [warn] missing keys: {len(missing)} (e.g. {missing[:3]})")
62
+ if unexpected:
63
+ print(f" [warn] unexpected keys: {len(unexpected)} (e.g. {unexpected[:3]})")
64
+ model.to(device).eval()
65
+ step = ckpt.get("step", "?")
66
+ print(f"loaded {ckpt_path} | step {step} | vocab {cfg.vocab_size} | "
67
+ f"block_size {cfg.block_size}")
68
+ return model, cfg
69
+
70
+
71
+ # --------------------------------------------------------------------------- #
72
+ # sampling -- tuned for the oscillator stack (no KV cache, sanitised logits)
73
+ # --------------------------------------------------------------------------- #
74
+ def _filter_logits(logits, top_k, top_p):
75
+ """Apply top-k then nucleus (top-p) filtering in place; returns logits."""
76
+ if top_k:
77
+ k = min(top_k, logits.size(-1))
78
+ kth = torch.topk(logits, k).values[..., -1, None]
79
+ logits = logits.masked_fill(logits < kth, float("-inf"))
80
+ if top_p and 0.0 < top_p < 1.0:
81
+ sorted_logits, sorted_idx = torch.sort(logits, descending=True, dim=-1)
82
+ cum = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
83
+ remove = cum > top_p
84
+ remove[..., 1:] = remove[..., :-1].clone() # keep the first over-threshold token
85
+ remove[..., 0] = False
86
+ scatter_remove = remove.scatter(-1, sorted_idx, remove)
87
+ logits = logits.masked_fill(scatter_remove, float("-inf"))
88
+ return logits
89
+
90
+
91
+ @torch.no_grad()
92
+ def generate(model, cfg, idx, n_new, temperature=0.9, top_k=40, top_p=0.95,
93
+ repetition_penalty=1.1, stop_ids=None, device="cpu", use_cache=True):
94
+ """Autoregressive sampling. Mirrors model.generate()'s NaN/inf sanitisation
95
+ (long oscillator rollouts can emit +/-inf logits whose softmax -> NaN ->
96
+ multinomial device-side assert) and adds top-p + repetition penalty.
97
+
98
+ Uses the attention KV cache by default: the prompt is prefilled once, then
99
+ each new token is a single-position forward. Once the cache reaches
100
+ block_size we re-prime with a windowed recompute (absolute RoPE positions
101
+ must stay within max_position_embeddings)."""
102
+ model.eval()
103
+ stop_ids = set(stop_ids or [])
104
+ past = None
105
+ for _ in range(n_new):
106
+ if use_cache and past is not None and past[0][0].size(2) < cfg.block_size:
107
+ logits, _, _, past = model(idx[:, -1:], past_key_values=past, use_cache=True)
108
+ elif use_cache:
109
+ logits, _, _, past = model(idx[:, -cfg.block_size:], use_cache=True)
110
+ else:
111
+ logits, _, _ = model(idx[:, -cfg.block_size:])
112
+ lg = torch.nan_to_num(logits[:, -1, :].float(),
113
+ nan=0.0, posinf=1e4, neginf=-1e4)
114
+
115
+ # repetition penalty over tokens already in the context
116
+ if repetition_penalty and repetition_penalty != 1.0:
117
+ for b in range(lg.size(0)):
118
+ seen = torch.unique(idx[b])
119
+ vals = lg[b, seen]
120
+ lg[b, seen] = torch.where(vals > 0, vals / repetition_penalty,
121
+ vals * repetition_penalty)
122
+
123
+ lg = lg / max(temperature, 1e-6)
124
+ lg = _filter_logits(lg, top_k, top_p)
125
+ probs = torch.softmax(lg, dim=-1)
126
+
127
+ if not torch.isfinite(probs).all() or float(probs.sum()) <= 0.0:
128
+ nxt = torch.argmax(torch.nan_to_num(lg, neginf=-1e4), dim=-1, keepdim=True)
129
+ else:
130
+ nxt = torch.multinomial(probs, 1)
131
+
132
+ idx = torch.cat([idx, nxt], dim=1)
133
+ if stop_ids and int(nxt[0]) in stop_ids and idx.size(0) == 1:
134
+ break
135
+ return idx
136
+
137
+
138
+ # --------------------------------------------------------------------------- #
139
+ # self-speculative decoding (DeepSpec draft+verify, drafted by the MTP heads)
140
+ # --------------------------------------------------------------------------- #
141
+ @torch.no_grad()
142
+ def generate_speculative(model, cfg, idx, n_new, stop_ids=None, verbose=False, trace=None):
143
+ """Greedy self-speculative decoding (batch size 1).
144
+
145
+ DeepSpec-style draft+verify, but the model drafts on ITSELF: the `mtp_layers`
146
+ MTP heads propose the next `mtp_layers` future tokens from one hidden state,
147
+ and the main head verifies them in a single parallel forward, accepting the
148
+ longest correct prefix. Output is BIT-IDENTICAL to greedy generate() (top_k=1)
149
+ -- speculation only changes the number of forwards, never the tokens.
150
+
151
+ Per cycle (one forward over committed seq + the pending drafts):
152
+ base = index of the last committed token; main_logits[base+j] is the
153
+ verifier's genuine token AFTER draft j. We accept draft j while it matches,
154
+ commit one correction token at the first mismatch, then read the next round's
155
+ drafts from the MTP heads at the deepest still-valid position (base+a)."""
156
+ assert idx.size(0) == 1, "speculative decoder is batch-1"
157
+ model.eval()
158
+ if model.mtp_heads is None:
159
+ raise ValueError("speculative decoding needs MTP heads (use_mtp=True).")
160
+ stop_ids = set(stop_ids or [])
161
+ K = len(model.mtp_heads)
162
+ start_len = idx.size(1)
163
+ forwards = 0
164
+
165
+ def crop(seq): # honor block_size context window
166
+ return seq[:, -cfg.block_size:]
167
+
168
+ # bootstrap: one forward over the prompt -> genuine next token + K drafts
169
+ main_logits, mtp = model.forward_drafts(crop(idx)); forwards += 1
170
+ g = main_logits[:, -1].argmax(-1, keepdim=True)
171
+ idx = torch.cat([idx, g], dim=1)
172
+ drafts = [mtp[k][:, -1].argmax(-1, keepdim=True) for k in range(K)] # guesses after g
173
+
174
+ while idx.size(1) - start_len < n_new:
175
+ if stop_ids and int(idx[0, -1]) in stop_ids:
176
+ break
177
+ seq = torch.cat([idx] + drafts, dim=1) # committed + K speculative tokens
178
+ base = idx.size(1) - 1 # position of the last committed token
179
+ cseq = crop(seq)
180
+ off = cseq.size(1) - seq.size(1) # crop shift (<=0) to realign `base`
181
+ main_logits, mtp = model.forward_drafts(cseq); forwards += 1
182
+ b = base + off
183
+
184
+ # verify: accept draft[j] while it matches the verifier's token after it
185
+ a = 0
186
+ for j in range(K):
187
+ v = main_logits[:, b + j].argmax(-1, keepdim=True) # genuine token after draft j-1 / g
188
+ if int(v) == int(drafts[j]):
189
+ a += 1
190
+ else:
191
+ break
192
+ accepted = drafts[:a]
193
+ correction = main_logits[:, b + a].argmax(-1, keepdim=True) # genuine token at first miss
194
+ new_toks = accepted + [correction]
195
+ idx = torch.cat([idx] + new_toks, dim=1)
196
+ if trace is not None: # per-cycle record for the visualizer
197
+ trace.append({"drafts": [int(d) for d in drafts], "accepted": a,
198
+ "correction": int(correction)})
199
+
200
+ # next drafts from the deepest still-valid hidden (position b+a holds the last
201
+ # accepted token / g, whose MTP heads predict the tokens AFTER `correction`).
202
+ drafts = [mtp[k][:, b + a].argmax(-1, keepdim=True) for k in range(K)]
203
+
204
+ if stop_ids and any(int(t) in stop_ids for t in new_toks):
205
+ break
206
+
207
+ idx = idx[:, :start_len + n_new]
208
+ if verbose:
209
+ produced = idx.size(1) - start_len
210
+ print(f" [spec] {produced} tokens in {forwards} forwards "
211
+ f"({produced / max(forwards,1):.2f} tok/forward; "
212
+ f"vs 1.00 for plain greedy)")
213
+ return idx
214
+
215
+
216
+ # --------------------------------------------------------------------------- #
217
+ # prompt helpers
218
+ # --------------------------------------------------------------------------- #
219
+ def build_prompt(tok, text, chat, system):
220
+ """Return token ids for the prompt. In chat mode, wrap in ChatML so the
221
+ model is cued to produce an assistant turn (only useful if it was trained
222
+ on that framing -- harmless otherwise)."""
223
+ if chat:
224
+ parts = []
225
+ if system:
226
+ parts.append(f"<|im_start|><|system|>\n{system}<|im_end|>\n")
227
+ parts.append(f"<|im_start|><|user|>\n{text}<|im_end|>\n")
228
+ parts.append("<|im_start|><|assistant|>\n")
229
+ text = "".join(parts)
230
+ ids = tok.encode(text, add_special_tokens=False)
231
+ return ids, text
232
+
233
+
234
+ def resolve_stop_ids(tok, chat):
235
+ """Collect ids of end-of-turn / end-of-text markers that exist in the vocab."""
236
+ names = ["<eos>", "<|endoftext|>"]
237
+ if chat:
238
+ names.insert(0, "<|im_end|>")
239
+ vocab = tok.get_vocab()
240
+ return [vocab[n] for n in names if n in vocab]
241
+
242
+
243
+ # --------------------------------------------------------------------------- #
244
+ # main
245
+ # --------------------------------------------------------------------------- #
246
+ def run_once(model, cfg, tok, args, device, prompt_text):
247
+ stop_ids = resolve_stop_ids(tok, args.chat)
248
+ ids, shown = build_prompt(tok, prompt_text, args.chat, args.system)
249
+ x = torch.tensor([ids], dtype=torch.long, device=device)
250
+ if args.speculative:
251
+ # greedy self-speculative decoding (MTP heads draft, main head verifies);
252
+ # output equals greedy decoding, only faster. Ignores sampling knobs.
253
+ out = generate_speculative(model, cfg, x, args.max_new_tokens,
254
+ stop_ids=stop_ids, verbose=args.spec_stats)
255
+ else:
256
+ out = generate(model, cfg, x, args.max_new_tokens,
257
+ temperature=args.temperature, top_k=args.top_k,
258
+ top_p=args.top_p, repetition_penalty=args.repetition_penalty,
259
+ stop_ids=stop_ids, device=device, use_cache=not args.no_cache)
260
+ gen_ids = out[0, len(ids):].tolist()
261
+ completion = tok.decode(gen_ids, skip_special_tokens=not args.show_special)
262
+ if args.echo_prompt:
263
+ print(shown, end="")
264
+ print(completion)
265
+ return completion
266
+
267
+
268
+ def main():
269
+ p = argparse.ArgumentParser(description="Quazimoto-LM generation harness")
270
+ p.add_argument("--ckpt", default=os.path.join(PKG_DIR, "chkpt", "quazimoto.pt"),
271
+ help="checkpoint .pt saved by train.py")
272
+ p.add_argument("--tok_dir", default=PKG_DIR, help="dir holding tokenizer.json")
273
+ p.add_argument("--prompt", default="the quazimoto oscillator")
274
+ p.add_argument("--max_new_tokens", type=int, default=200)
275
+ p.add_argument("--temperature", type=float, default=0.9)
276
+ p.add_argument("--top_k", type=int, default=40, help="0 to disable")
277
+ p.add_argument("--top_p", type=float, default=0.95, help="1.0 to disable")
278
+ p.add_argument("--repetition_penalty", type=float, default=1.1, help="1.0 to disable")
279
+ p.add_argument("--seed", type=int, default=None)
280
+ p.add_argument("--chat", action="store_true", help="wrap prompt in ChatML framing")
281
+ p.add_argument("--system", default="", help="system prompt (chat mode only)")
282
+ p.add_argument("--echo_prompt", action="store_true", help="print the prompt before output")
283
+ p.add_argument("--show_special", action="store_true", help="don't strip special tokens")
284
+ p.add_argument("--interactive", action="store_true", help="REPL: read prompts from stdin")
285
+ p.add_argument("--no_cache", action="store_true",
286
+ help="disable the KV cache (full recompute each step)")
287
+ p.add_argument("--speculative", action="store_true",
288
+ help="greedy self-speculative decoding via MTP draft heads (DeepSpec-style)")
289
+ p.add_argument("--spec_stats", action="store_true",
290
+ help="print tokens/forward acceptance stats for speculative decoding")
291
+ p.add_argument("--device", default="cpu" if torch.cuda.is_available() else "cpu")
292
+ args = p.parse_args()
293
+
294
+ if args.top_k == 0:
295
+ args.top_k = None
296
+ if args.seed is not None:
297
+ torch.manual_seed(args.seed)
298
+
299
+ device = args.device
300
+ tok = load_tokenizer(args.tok_dir)
301
+ model, cfg = load_model(args.ckpt, device)
302
+
303
+ if cfg.vocab_size != tok.vocab_size:
304
+ print(f" [warn] model vocab {cfg.vocab_size} != tokenizer vocab "
305
+ f"{tok.vocab_size}; decode may be misaligned.")
306
+
307
+ if args.interactive:
308
+ print("interactive mode -- type a prompt, Ctrl-C / empty line + EOF to quit.\n")
309
+ try:
310
+ while True:
311
+ try:
312
+ line = input(">>> ").strip()
313
+ except EOFError:
314
+ break
315
+ if not line:
316
+ continue
317
+ run_once(model, cfg, tok, args, device, line)
318
+ print()
319
+ except KeyboardInterrupt:
320
+ print("\nbye.")
321
+ else:
322
+ run_once(model, cfg, tok, args, device, args.prompt)
323
+
324
+
325
+ if __name__ == "__main__":
326
+ main()
healthcheck.py ADDED
@@ -0,0 +1,258 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ healthcheck.py -- checkpoint health report for Quazimoto-LM.
3
+
4
+ Auto-finds the checkpoint (same rule as train.py --resume: the --out/--ckpt file
5
+ if present, else the newest *.pt in the folder), then reports:
6
+ 1. load + global tensor health (NaN/Inf, dtypes, param count, step)
7
+ 2. per-(grouped)-layer weight stats with dead/exploded flags
8
+ 3. family-trait layer status -- each trait's gates (tanh) so you can see which
9
+ blocks have actually woken up, plus their key weight stats
10
+ 4. current perplexity on a sample of text (the local --data file, or a built-in
11
+ fallback string), with the random-vocab baseline for reference.
12
+
13
+ Usage:
14
+ python healthcheck.py # auto-find in ./chkpt
15
+ python healthcheck.py --ckpt path/to.pt
16
+ python healthcheck.py --data corpus.txt # measure PPL on your own text
17
+ """
18
+
19
+ import argparse
20
+ import math
21
+ import os
22
+ import re
23
+ import sys
24
+ from collections import Counter
25
+
26
+ import torch
27
+
28
+ from model import QuazimotoLM, QuazimotoConfig
29
+ from family import (HRMRefinementBlock, JEPAPredictorBlock, MoESwiGLU, MTPHead,
30
+ EngramRing, PhaseAttentionRing, RingControllerBank, RingSpecialists)
31
+
32
+ PKG_DIR = os.path.dirname(os.path.abspath(__file__))
33
+
34
+
35
+ # --------------------------------------------------------------------------- #
36
+ # discovery / loading
37
+ # --------------------------------------------------------------------------- #
38
+ def find_ckpt(path):
39
+ """Prefer an exact file; else newest *.pt in its folder (or ./chkpt)."""
40
+ if path and os.path.isfile(path):
41
+ return path
42
+ folder = os.path.dirname(path) if path else os.path.join(PKG_DIR, "chkpt")
43
+ folder = folder or "."
44
+ if not os.path.isdir(folder):
45
+ return None
46
+ pts = [os.path.join(folder, f) for f in os.listdir(folder) if f.endswith(".pt")]
47
+ return max(pts, key=os.path.getmtime) if pts else None
48
+
49
+
50
+ def load_tokenizer(tok_dir):
51
+ sys.path.insert(0, tok_dir)
52
+ from spike_tokenizer import SpikeTokenizer
53
+ return SpikeTokenizer(vocab_file=os.path.join(tok_dir, "tokenizer.json"))
54
+
55
+
56
+ def section(title):
57
+ print("\n" + "=" * 72 + f"\n{title}\n" + "=" * 72)
58
+
59
+
60
+ def fmt_gate(t):
61
+ return "[" + ", ".join(f"{x:+.3f}" for x in t) + "]"
62
+
63
+
64
+ # --------------------------------------------------------------------------- #
65
+ # 1. global tensor health
66
+ # --------------------------------------------------------------------------- #
67
+ def report_load(ckpt, path):
68
+ section("1. LOAD + GLOBAL TENSOR HEALTH")
69
+ print(f"checkpoint : {path}")
70
+ print(f"keys : {list(ckpt.keys())}")
71
+ print(f"step : {ckpt.get('step')} | saved vocab_size: {ckpt.get('vocab_size')}")
72
+ print(f"optimizer : {'present (resumable)' if 'optim' in ckpt else 'absent'}")
73
+ sd = ckpt["model"]
74
+ nan = inf = total = 0
75
+ dtypes = Counter()
76
+ for v in sd.values():
77
+ if not torch.is_tensor(v):
78
+ continue
79
+ total += v.numel(); dtypes[str(v.dtype)] += 1
80
+ nan += torch.isnan(v).sum().item(); inf += torch.isinf(v).sum().item()
81
+ print(f"tensors : {len(sd)} | params: {total/1e6:.1f}M | dtypes: {dict(dtypes)}")
82
+ flag = "OK" if (nan == 0 and inf == 0) else "*** PROBLEM ***"
83
+ print(f"NaNs : {nan} | Infs: {inf} -> {flag}")
84
+ return sd
85
+
86
+
87
+ # --------------------------------------------------------------------------- #
88
+ # 2. per-layer weight stats
89
+ # --------------------------------------------------------------------------- #
90
+ def report_layers(sd):
91
+ section("2. PER-LAYER WEIGHT HEALTH (grouped over repeated layers)")
92
+ buckets = {}
93
+ for k, v in sd.items():
94
+ if not torch.is_tensor(v) or v.dtype != torch.float32:
95
+ continue
96
+ base = re.sub(r"\.\d+\.", ".N.", k)
97
+ buckets.setdefault(base, []).append(v)
98
+ print(f"{'param group':46s} {'std':>8} {'absmax':>9} {'mean':>9} flag")
99
+ for base, vs in buckets.items():
100
+ if "ring_bank" in base or "specialists" in base:
101
+ continue # shown in the traits section
102
+ s = torch.cat([x.flatten() for x in vs]).float()
103
+ if s.numel() < 2:
104
+ continue # scalars -> traits section
105
+ flag = ""
106
+ if s.std() < 1e-4 and s.abs().max() < 1e-4: flag = "DEAD?" # all ~0 (not just constant)
107
+ if s.abs().max() > 50: flag = "EXPLODED?"
108
+ print(f"{base:46s} {s.std():8.4f} {s.abs().max():9.3f} {s.mean():+9.4f} {flag}")
109
+
110
+
111
+ # --------------------------------------------------------------------------- #
112
+ # 3. family-trait layer status
113
+ # --------------------------------------------------------------------------- #
114
+ def _wstat(t):
115
+ t = t.float()
116
+ return f"std {t.std():.4f} absmax {t.abs().max():.3f}" if t.numel() > 1 else f"val {t.item():+.4f}"
117
+
118
+
119
+ def report_traits(model, cfg):
120
+ section("3. FAMILY-TRAIT LAYER STATUS (tanh-gate ~0 => still a no-op)")
121
+ any_trait = False
122
+
123
+ # --- per-layer Quazimoto oscillator block gate + interstitial/specialist rings ---
124
+ for li, layer in enumerate(model.layers):
125
+ q = layer.quaz
126
+ tag = f"layer{li}.quaz"
127
+ print(f"\n[{tag}] oscillator block gate tanh = {torch.tanh(q.gate).item():+.4f} "
128
+ f"| block_gain {_wstat(q.block_gain)} | alpha {_wstat(q.alpha)}")
129
+ if hasattr(q, "phase_seed_gate"):
130
+ print(f" fractal phase-seed gate : {torch.tanh(q.phase_seed_gate).item():+.4f} "
131
+ f"(0 => seed unused)")
132
+ if getattr(q, "use_rings", False):
133
+ ar = [torch.tanh(r.gate).item() for r in q.attn_rings]
134
+ er = [torch.tanh(r.scale).item() for r in q.engram_rings]
135
+ print(f" phase-attn ring gates : {fmt_gate(ar)}")
136
+ print(f" engram ring scales : {fmt_gate(er)}")
137
+ if getattr(q, "use_ring_specialists", False):
138
+ sc = [torch.tanh(b.scale).item() for b in q.specialists]
139
+ act = [b.active.mean().item() for b in q.specialists]
140
+ filled = [float(b.store_out.abs().sum() > 0) for b in q.specialists]
141
+ print(f" specialist out-gates : {fmt_gate(sc)}")
142
+ print(f" specialist mean 'active': {fmt_gate(act)}")
143
+ print(f" specialist store filled : {fmt_gate(filled)} (0 until written at inference)")
144
+ if li == 0 and (getattr(q, 'use_rings', False) or getattr(q, 'use_ring_specialists', False)):
145
+ any_trait = True
146
+
147
+ # --- trunk-level traits ---
148
+ if model.hrm is not None:
149
+ any_trait = True
150
+ g = torch.tanh(model.hrm.gate)
151
+ print(f"\n[HRM] gate tanh = {fmt_gate(g.tolist())} "
152
+ f"({'OPEN' if g.abs().mean() > 0.01 else 'no-op'}) | "
153
+ f"z0 {_wstat(model.hrm.z0)} | down {_wstat(model.hrm.down.weight)}")
154
+ if model.moe is not None:
155
+ any_trait = True
156
+ print(f"\n[MoE] router {_wstat(model.moe.router.weight)} | "
157
+ f"experts={model.moe.n_routed} shared={model.moe.n_shared} top_k={model.moe.top_k}")
158
+ if model.mtp_heads is not None:
159
+ any_trait = True
160
+ norms = [h.proj.weight.norm().item() for h in model.mtp_heads]
161
+ print(f"\n[MTP] {len(model.mtp_heads)} draft heads | proj-weight norms "
162
+ f"{fmt_gate(norms)} (0 => head still identity / untrained)")
163
+ if model.jepa is not None:
164
+ any_trait = True
165
+ g = torch.tanh(model.jepa.gate)
166
+ print(f"\n[JEPA] gate tanh = {fmt_gate(g.tolist())} (train-only aux; never runs in generate)")
167
+ if model.ring_bank is not None:
168
+ any_trait = True
169
+ decs = [c.dec.weight.norm().item() for c in model.ring_bank.controllers]
170
+ print(f"\n[Controllers] {len(model.ring_bank.controllers)} ring controllers | "
171
+ f"decoder-weight norms {fmt_gate(decs)} (0 => no-op self-organizer)")
172
+ if getattr(cfg, "use_fractal_phase_seed", False):
173
+ any_trait = True
174
+ gates = [torch.tanh(layer.quaz.phase_seed_gate).item() for layer in model.layers]
175
+ mag = sum(abs(g) for g in gates) / max(len(gates), 1)
176
+ tbl = os.path.join(PKG_DIR, "fractal_phase.pt")
177
+ if os.path.exists(tbl):
178
+ d = torch.load(tbl, map_location="cpu", weights_only=False)
179
+ ok = tuple(d.get("phases").shape) == (cfg.vocab_size, cfg.n_osc)
180
+ tinfo = f"{d.get('mode','?')} table {'MATCHES' if ok else '** SHAPE MISMATCH **'} " \
181
+ f"{tuple(d.get('phases').shape)}"
182
+ else:
183
+ tinfo = "no fractal_phase.pt found -> model uses FLAT Halton fallback"
184
+ print(f"\n[FractalSeed] per-layer phase-seed gate tanh {fmt_gate(gates)}")
185
+ print(f" status: {'OPEN' if mag > 0.01 else 'no-op (gates ~0)'} | {tinfo}")
186
+
187
+ if not any_trait:
188
+ print("\n(no opt-in family traits enabled in this checkpoint)")
189
+
190
+
191
+ # --------------------------------------------------------------------------- #
192
+ # 4. perplexity
193
+ # --------------------------------------------------------------------------- #
194
+ @torch.no_grad()
195
+ def report_ppl(model, cfg, tok, data_path, device):
196
+ section("4. PERPLEXITY")
197
+ if data_path and os.path.exists(data_path):
198
+ with open(data_path, "r", encoding="utf-8", errors="replace") as f:
199
+ text = f.read()[:200000]
200
+ src = data_path
201
+ else:
202
+ text = ("the quazimoto oscillator learns to synchronize phases across rings. "
203
+ "coupled clocks fall into step when coupling is strong enough. ") * 50
204
+ src = "built-in fallback string"
205
+ ids = tok.encode(text, add_special_tokens=False)
206
+ ids = ids[: cfg.block_size * 8 + 1] # cap cost; a few windows is enough
207
+ if len(ids) < 2:
208
+ print("not enough tokens to score."); return
209
+ x = torch.tensor([ids], device=device)
210
+ losses = []
211
+ B = cfg.block_size
212
+ for s in range(0, x.size(1) - 1, B): # non-overlapping windows
213
+ xb = x[:, s:s + B]; yb = x[:, s + 1:s + 1 + B]
214
+ n = min(xb.size(1), yb.size(1))
215
+ if n < 1: break
216
+ _, loss, _ = model(xb[:, :n], yb[:, :n])
217
+ losses.append(loss.item())
218
+ ce = sum(losses) / len(losses)
219
+ print(f"text source : {src} ({len(ids)} tokens, {len(losses)} windows)")
220
+ print(f"CE loss : {ce:.4f}")
221
+ print(f"bits/token : {ce/math.log(2):.4f}")
222
+ print(f"perplexity : {math.exp(ce):.1f} (random-vocab baseline: {cfg.vocab_size})")
223
+
224
+
225
+ # --------------------------------------------------------------------------- #
226
+ def main():
227
+ p = argparse.ArgumentParser(description="Quazimoto-LM checkpoint health check")
228
+ p.add_argument("--ckpt", default="", help="checkpoint path (default: auto-find newest in ./chkpt)")
229
+ p.add_argument("--data", default="", help="text file to measure PPL on (default: built-in string)")
230
+ p.add_argument("--tok_dir", default=PKG_DIR)
231
+ p.add_argument("--device", default="cpu")
232
+ args = p.parse_args()
233
+
234
+ path = find_ckpt(args.ckpt)
235
+ if path is None:
236
+ print("No checkpoint found. Pass --ckpt or train one first.")
237
+ return
238
+
239
+ ckpt = torch.load(path, map_location=args.device, weights_only=False)
240
+ sd = report_load(ckpt, path)
241
+ report_layers(sd)
242
+
243
+ cfg = QuazimotoConfig(**ckpt["family_config"])
244
+ model = QuazimotoLM(cfg)
245
+ miss, unexp = model.load_state_dict(sd, strict=False)
246
+ model.to(args.device).eval()
247
+ if miss or unexp:
248
+ print(f"\n[load] missing {len(miss)} / unexpected {len(unexp)} keys "
249
+ f"(arch drift vs this code; e.g. {(miss or unexp)[:2]})")
250
+
251
+ tok = load_tokenizer(args.tok_dir)
252
+ report_traits(model, cfg)
253
+ report_ppl(model, cfg, tok, args.data, args.device)
254
+ print()
255
+
256
+
257
+ if __name__ == "__main__":
258
+ main()
instrument.py ADDED
@@ -0,0 +1,84 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ instrument.py -- zero-overhead-when-off capture hooks for the live visualizer.
3
+
4
+ A single global Recorder is consulted by the model's hot paths. When no recorder
5
+ is active (the default), every call is a single `is None` check, so training and
6
+ normal generation are unaffected. visualize.py activates a recorder, runs a
7
+ forward per generated token, and reads back per-token state.
8
+ """
9
+
10
+ _REC = None
11
+
12
+
13
+ def get_rec():
14
+ return _REC
15
+
16
+
17
+ def set_rec(r):
18
+ global _REC
19
+ _REC = r
20
+
21
+
22
+ def _r(x, nd=4):
23
+ return round(float(x), nd)
24
+
25
+
26
+ class Recorder:
27
+ """Collects one frame per generated token. Modules append in execution order,
28
+ so list position == layer index (attention then quaz, per layer)."""
29
+
30
+ def __init__(self, phase_layer=None, attn_layer=None):
31
+ self.frames = []
32
+ self.cur = None
33
+ self.enabled = False
34
+ self.phase_layer = phase_layer # which layer's individual phases to keep
35
+ self.attn_layer = attn_layer # which layer's attention map to keep
36
+ self._spec_tmp = []
37
+
38
+ # ---- frame lifecycle (driven by visualize.py) ----
39
+ def begin(self):
40
+ self.cur = {"rings": [], "phases": None, "attn": None,
41
+ "spec": [], "quaz_norm": [], "traits": {}}
42
+ self._spec_tmp = []
43
+ self.enabled = True
44
+
45
+ def end(self, **meta):
46
+ self.enabled = False
47
+ if self.cur is not None:
48
+ self.cur.update(meta)
49
+ self.frames.append(self.cur)
50
+ self.cur = None
51
+
52
+ # ---- module callbacks (guarded by `enabled` at the call site) ----
53
+ def log_ring(self, R, psi, phases):
54
+ li = len(self.cur["rings"])
55
+ self.cur["rings"].append({"R": [_r(x, 4) for x in R],
56
+ "psi": [_r(x, 4) for x in psi]})
57
+ if phases is not None and (self.phase_layer is None or li == self.phase_layer):
58
+ self.cur["phases"] = {"layer": li, "theta": [_r(x, 3) for x in phases]}
59
+
60
+ def log_traj(self, traj):
61
+ """Attach the tip growth trajectory (list over growth steps, each a flattened
62
+ [N*pos_dim] snapshot) to THIS layer's colony frame (phase_layer only). Lets the
63
+ 3D view draw each tip's hypha extending across the growth steps."""
64
+ ph = self.cur.get("phases")
65
+ if ph is not None and ph.get("layer") == len(self.cur["rings"]) - 1:
66
+ ph["traj"] = [[_r(x, 3) for x in step] for step in traj]
67
+
68
+ def log_attn(self, layer_idx, w):
69
+ if self.attn_layer is None or layer_idx == self.attn_layer:
70
+ self.cur["attn"] = {"layer": layer_idx, "w": [_r(x, 4) for x in w]}
71
+
72
+ def log_quaz_norm(self, n):
73
+ self.cur["quaz_norm"].append(_r(n, 4))
74
+
75
+ def push_spec(self, route): # one ring's [n_spec] route weights
76
+ self._spec_tmp.append([_r(x, 4) for x in route])
77
+
78
+ def flush_spec(self): # called once per quaz block (a layer)
79
+ if self._spec_tmp:
80
+ self.cur["spec"].append(self._spec_tmp)
81
+ self._spec_tmp = []
82
+
83
+ def log_trait(self, name, val):
84
+ self.cur["traits"][name] = _r(val, 4)
model.py ADDED
@@ -0,0 +1,773 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Quazimoto-LM — a language model whose channel-mixing block is a bank of
3
+ coupled phase oscillators (Kuramoto dynamics) arranged in concentric rings.
4
+
5
+ Backbone: standard causal transformer (token mixing = attention, tied LM head).
6
+ Novelty: the per-layer MLP is replaced by a QuazimotoBlock that maps the hidden
7
+ state to N oscillator phases on rings [4, 8, 16, 32, ...], runs a few
8
+ differentiable Euler steps of structured Kuramoto coupling with a
9
+ learnable frustration alpha, then reads out [cos, sin] of the phases.
10
+
11
+ Set ring coupling to dense + alpha=0 and it degenerates toward AKOrN-style
12
+ oscillatory neurons; the ring structure + hierarchy-to-center is the Quazimoto bit.
13
+ """
14
+
15
+ from dataclasses import dataclass, field
16
+ import math
17
+ import os
18
+ import torch
19
+ import torch.nn as nn
20
+ import torch.nn.functional as F
21
+
22
+ _PKG_DIR = os.path.dirname(os.path.abspath(__file__))
23
+
24
+ # family traits (DERF soft_clamp, RMSNorm, and the opt-in refinement/aux blocks)
25
+ from family import (soft_clamp, RMSNorm, HRMRefinementBlock, MoESwiGLU,
26
+ MTPHead, JEPAPredictorBlock, PhaseAttentionRing, EngramRing,
27
+ RingControllerBank, RingSpecialists)
28
+ import instrument as _viz # live-visualizer capture hooks (no-op unless a recorder is active)
29
+ from mycel import MycelBlock, MycelStations # Neighbour-Sensing growth mixer
30
+
31
+
32
+ @dataclass
33
+ class QuazimotoConfig:
34
+ vocab_size: int = 32000
35
+ n_layer: int = 10
36
+ n_head: int = 12
37
+ d_model: int = 768
38
+ block_size: int = 1024
39
+ dropout: float = 0.0
40
+ # Quazimoto oscillator bank
41
+ ring_sizes: tuple = (4, 8, 16, 32, 64, 128, 256) # (unused by the mycelium mixer)
42
+ osc_steps: int = 4 # differentiable Kuramoto Euler steps
43
+ osc_dt: float = 0.5
44
+ readout_mult: int = 3 # MLP expansion on the readout
45
+ # Neighbour-Sensing (mycelium) mixer: N growing tips in a bounded latent region.
46
+ n_tips: int = 96 # hyphal tips per token (analog of n_osc)
47
+ mycel_pos_dim: int = 3 # latent dimensionality tips grow in
48
+ mycel_field_centers: int = 16 # low-rank density-field sample sites (O(N*F) not O(N^2))
49
+ mycel_n_stations: int = 16 # trait stations spread through the colony
50
+ growth_steps: int = 3 # differentiable growth iterations (each retains ~5 [B,T,N,F])
51
+ growth_dt: float = 0.3
52
+ # Mandelbrot phase seeding: seed the N oscillator phases from each token's
53
+ # fractal orbit angles (frozen, parameter-free), added through a zero-init gate.
54
+ use_fractal_phase_seed: bool = False
55
+ # family-trait discipline
56
+ osc_bound: float = 10.0 # DERF soft_clamp bound on block output (BPTT-stability)
57
+ gate_init_open: float = 0.9 # tanh-gate opening at init (Quazimoto is the main mixer,
58
+ # so it starts OPEN; set 0.0 for a pure no-op refinement)
59
+ # interstitial rings: between each pair of oscillator rings, insert a phase-attention
60
+ # ring + an engram ring that ABSORB info and INJECT it into the two neighbor rings.
61
+ use_rings: bool = False
62
+ ring_attn_heads: int = 4
63
+ ring_attn_head_dim: int = 16
64
+ ring_engram_compress: int = 32
65
+ ring_engram_heads: int = 2
66
+ ring_engram_table: int = 2048
67
+ ring_engram_ngram: int = 3
68
+ # per-ring memory specialists: a small MoE of `ring_n_specialists` mini experts
69
+ # PER oscillator ring, each holding test-time-writable input/output stores that
70
+ # act as an addressable context memory at inference (see family.RingSpecialists).
71
+ use_ring_specialists: bool = False
72
+ ring_n_specialists: int = 7
73
+ ring_spec_key_dim: int = 32
74
+ ring_spec_slot_dim: int = 64
75
+ ring_spec_top_k: int = 2
76
+ ring_spec_write_lr: float = 0.1
77
+ # per-ring self-organizing controllers (~1M total: one tiny net per oscillator
78
+ # ring, shared across layers, weights self-optimize by a predictive/free-energy rule)
79
+ use_ring_controllers: bool = False
80
+ ring_ctrl_feat: int = 384 # fast-weight predictor width (dominates the ~150k/ring)
81
+ ring_ctrl_local_lr: float = 0.01 # delta-rule (surprise-minimization) step size
82
+ # opt-in family trait modules (all safe no-ops at init)
83
+ use_hrm: bool = False
84
+ hrm_steps: int = 3
85
+ hrm_dim: int = 256
86
+ hrm_gate_init: float = 0.1 # HRM gates start OPEN (random initial state needs a path out)
87
+ use_moe: bool = False
88
+ moe_intermediate: int = 768
89
+ moe_n_routed: int = 4
90
+ moe_n_shared: int = 1
91
+ moe_top_k: int = 2
92
+ use_mtp: bool = False
93
+ mtp_layers: int = 4 # depth of MTP draft heads (predict +1..+mtp_layers); used for spec-decode
94
+ mtp_loss_weight: float = 0.3
95
+ use_jepa: bool = False
96
+ jepa_pred_dim: int = 256
97
+ jepa_horizon: int = 1
98
+ jepa_loss_weight: float = 0.1
99
+ # ---- attention: ported from model_v2.MLADerfXSAAttention ----
100
+ head_dim: int = 64
101
+ qk_rope_head_dim: int = 32 # partial RoPE: rope part of each head
102
+ nope_head_dim: int = 32 # no-pos part (sum == head_dim)
103
+ num_key_value_heads: int = 4 # GQA (== n_head for full MHA)
104
+ q_lora_rank: int = 384 # MLA low-rank q projection
105
+ o_lora_rank: int = 384 # MLA low-rank output projection
106
+ use_qk_norm: bool = True # per-head RMSNorm on Q/K before RoPE
107
+ use_derf: bool = False # erf attention instead of softmax (ablate)
108
+ use_xsa: bool = False # value-subspace removal (ablate)
109
+ rope_theta: float = 10000.0
110
+ max_position_embeddings: int = 4096
111
+ rms_norm_eps: float = 1e-6
112
+ initializer_range: float = 0.02
113
+ # ---- model-level v2 features ----
114
+ zloss_coef: float = 1e-4 # z-loss on lm logits
115
+ use_value_embed: bool = False # per-block value-embedding residual (zero-init gate)
116
+ use_hyper_connections: bool = False
117
+ hc_mult: int = 2
118
+
119
+ @property
120
+ def n_osc(self):
121
+ return self.n_tips * self.mycel_pos_dim # fractal seeds tip POSITIONS here
122
+
123
+
124
+ def _ring_index(ring_sizes):
125
+ """Return a LongTensor of length N mapping each oscillator -> its ring id."""
126
+ idx = []
127
+ for r, n in enumerate(ring_sizes):
128
+ idx += [r] * n
129
+ return torch.tensor(idx, dtype=torch.long)
130
+
131
+
132
+ class QuazimotoBlock(nn.Module):
133
+ """Oscillatory channel-mixing block: hidden state -> ring phases -> Kuramoto -> readout."""
134
+
135
+ def __init__(self, cfg: QuazimotoConfig):
136
+ super().__init__()
137
+ self.cfg = cfg
138
+ N = cfg.n_osc
139
+ R = len(cfg.ring_sizes)
140
+ self.N, self.R = N, R
141
+
142
+ self.norm = RMSNorm(cfg.d_model)
143
+ # project hidden -> initial phases and natural frequencies
144
+ self.to_theta = nn.Linear(cfg.d_model, N)
145
+ self.to_omega = nn.Linear(cfg.d_model, N)
146
+
147
+ # structured, LEARNABLE coupling: an RxR gain between ring blocks.
148
+ # Block-constant coupling lets us use per-ring mean fields (order
149
+ # parameters) -> O(N*R) instead of the O(N^2) pairwise sum.
150
+ rid = _ring_index(cfg.ring_sizes)
151
+ self.register_buffer("ring_id", rid)
152
+ onehot = F.one_hot(rid, R).float() # [N,R] ring membership
153
+ self.register_buffer("onehot", onehot)
154
+ self.block_gain = nn.Parameter(self._init_gain(R))
155
+ self.center_phase = nn.Parameter(torch.zeros(())) # the rotating center ball
156
+ self.center_omega = nn.Parameter(torch.tensor(0.3))
157
+ self.center_gain = nn.Parameter(torch.full((R,), 0.5)) # ring -> center pull
158
+ self.alpha = nn.Parameter(torch.zeros(R)) # per-ring frustration
159
+
160
+ # readout MLP on [cos theta, sin theta]
161
+ hidden = cfg.readout_mult * cfg.d_model
162
+ self.readout = nn.Sequential(
163
+ nn.Linear(2 * N, hidden),
164
+ nn.GELU(),
165
+ nn.Linear(hidden, cfg.d_model),
166
+ )
167
+ self.drop = nn.Dropout(cfg.dropout)
168
+
169
+ # family-trait gate: out = readout * tanh(gate). gate_init_open seeds it so the
170
+ # block is a warm-startable signal whose magnitude the optimizer controls; the
171
+ # content weights are nonzero so the gate always receives gradient.
172
+ go = math.atanh(min(cfg.gate_init_open, 0.9)) if cfg.gate_init_open > 0 else 0.0
173
+ self.gate = nn.Parameter(torch.full((1,), go))
174
+
175
+ # Mandelbrot phase-seed gate: theta <- to_theta(h) + tanh(seed_gate) * orbit.
176
+ # Zero-init -> no-op at start; opens only if the fractal prior helps.
177
+ if cfg.use_fractal_phase_seed:
178
+ self.phase_seed_gate = nn.Parameter(torch.zeros(1))
179
+
180
+ # std=0.02 init discipline (matches the family's linear-layer init)
181
+ for m in self.modules():
182
+ if isinstance(m, nn.Linear):
183
+ nn.init.normal_(m.weight, std=0.02)
184
+ if m.bias is not None:
185
+ nn.init.zeros_(m.bias)
186
+
187
+ # interstitial rings (instantiated AFTER the std-init loop). Oscillators are
188
+ # ordered ring-by-ring, so rings r and r+1 are the contiguous slice
189
+ # [off_r, off_{r+2}); the slot between them injects into exactly that block.
190
+ self.use_rings = cfg.use_rings
191
+ self.slot_bounds = []
192
+ if self.use_rings:
193
+ off = [0]
194
+ for n in cfg.ring_sizes:
195
+ off.append(off[-1] + n)
196
+ self.slot_bounds = [(off[r], off[r + 2]) for r in range(R - 1)]
197
+ self.attn_rings = nn.ModuleList([
198
+ PhaseAttentionRing(e - s, cfg.ring_attn_heads, cfg.ring_attn_head_dim, cfg.osc_bound)
199
+ for (s, e) in self.slot_bounds])
200
+ self.engram_rings = nn.ModuleList([
201
+ EngramRing(cfg.d_model, e - s, cfg.ring_engram_compress, cfg.ring_engram_heads,
202
+ cfg.ring_engram_table, cfg.ring_engram_ngram, cfg.osc_bound)
203
+ for (s, e) in self.slot_bounds])
204
+
205
+ # one memory-specialist MoE bank PER oscillator ring (injects into that ring)
206
+ self.use_ring_specialists = cfg.use_ring_specialists
207
+ if self.use_ring_specialists:
208
+ self.specialists = nn.ModuleList([
209
+ RingSpecialists(n, cfg.d_model, cfg.ring_n_specialists, cfg.ring_spec_key_dim,
210
+ cfg.ring_spec_slot_dim, cfg.ring_spec_top_k,
211
+ cfg.ring_spec_write_lr, cfg.osc_bound)
212
+ for n in cfg.ring_sizes])
213
+
214
+ @staticmethod
215
+ def _init_gain(R):
216
+ # within-ring coupling strong on the diagonal, weak neighbor coupling
217
+ g = torch.zeros(R, R)
218
+ for i in range(R):
219
+ g[i, i] = 1.0
220
+ if i + 1 < R:
221
+ g[i, i + 1] = g[i + 1, i] = 0.3
222
+ return g
223
+
224
+ def forward(self, x, ring_ctl=None, phase_seed=None):
225
+ cfg = self.cfg
226
+ B, T, _ = x.shape
227
+ h = self.norm(x)
228
+
229
+ theta = self.to_theta(h) # [B,T,N] initial phases
230
+ if phase_seed is not None: # Mandelbrot orbit phase prior (gated, no-op at init)
231
+ theta = theta + torch.tanh(self.phase_seed_gate) * phase_seed
232
+ omega = torch.tanh(self.to_omega(h)) # [B,T,N] natural frequencies (bounded)
233
+
234
+ # per-ring self-organizing controller: observe each ring (order-parameter
235
+ # magnitude/phase + mean freq), get no-op-at-init modulations of that ring's
236
+ # coupling / frustration / center-pull / injection.
237
+ d_coup = d_alpha = d_center = d_inj = None
238
+ if ring_ctl is not None:
239
+ with torch.no_grad():
240
+ c0, s0 = torch.cos(theta), torch.sin(theta)
241
+ cnt = self.onehot.sum(0).clamp(min=1.0) # [R] ring sizes
242
+ zc = (c0 @ self.onehot) / cnt # [B,T,R]
243
+ zs = (s0 @ self.onehot) / cnt
244
+ mag = torch.sqrt(zc ** 2 + zs ** 2 + 1e-8)
245
+ om = (omega @ self.onehot) / cnt
246
+ obs = torch.stack([mag.mean((0, 1)), (zc / mag).mean((0, 1)),
247
+ (zs / mag).mean((0, 1)), om.mean((0, 1))], dim=-1) # [R,4]
248
+ ctrl = ring_ctl(obs) # [R,4] (grad -> enc/dec only)
249
+ d_coup = torch.tanh(ctrl[:, 0]) # ring coupling scale (1 + .)
250
+ d_alpha = 0.5 * torch.tanh(ctrl[:, 1]) # frustration shift
251
+ d_center = torch.tanh(ctrl[:, 2]) # center-pull scale (1 + .)
252
+ d_inj = torch.tanh(ctrl[:, 3]) # injection scale (1 + .)
253
+
254
+ # interstitial rings absorb info ONCE (from the initial phases / hidden) and
255
+ # inject a constant drive into their two neighbor oscillator rings' phase update.
256
+ inject = 0.0
257
+ if self.use_rings:
258
+ inject = torch.zeros_like(theta)
259
+ for (s, e), ar, er in zip(self.slot_bounds, self.attn_rings, self.engram_rings):
260
+ contrib = ar(theta[..., s:e]) + er(h) # [B,T,e-s]
261
+ inject = inject + F.pad(contrib, (s, self.N - e)) # place at [s:e], sum overlaps
262
+ if d_inj is not None:
263
+ inject = inject * (1.0 + d_inj)[self.ring_id] # controller scales the drive
264
+
265
+ # memory specialists: each per-ring bank injects retrieved store_out into its
266
+ # ring's oscillators. Rings are contiguous and tile [0,N), so cat == full drive.
267
+ if self.use_ring_specialists:
268
+ spec_inj = torch.cat([bank(h) for bank in self.specialists], dim=-1) # [B,T,N]
269
+ inject = inject + spec_inj
270
+
271
+ rid = self.ring_id # [N]
272
+ gain = self.block_gain if d_coup is None else self.block_gain * (1.0 + d_coup).unsqueeze(1)
273
+ G = gain[rid] # [N,R] per-oscillator ring gains
274
+ alpha_vec = self.alpha if d_alpha is None else self.alpha + d_alpha
275
+ alpha_i = alpha_vec[rid] # [N]
276
+ center = self.center_phase # scalar phase, advances in time
277
+ cpull = self.center_gain if d_center is None else self.center_gain * (1.0 + d_center)
278
+ center_pull = cpull[rid] # [N]
279
+ invN = 1.0 / self.N
280
+
281
+ for s in range(cfg.osc_steps):
282
+ ctr = center + self.center_omega * (s * cfg.osc_dt)
283
+ c, sn = torch.cos(theta), torch.sin(theta) # [B,T,N]
284
+ # per-ring summed mean fields, then routed back through gains G
285
+ zc = (c @ self.onehot) @ G.t() # [B,T,N]
286
+ zs = (sn @ self.onehot) @ G.t() # [B,T,N]
287
+ A = alpha_i - theta # [B,T,N]
288
+ # (1/N) * sum_r G_ir * Im( e^{i(alpha_i - theta_i)} * Zsum_r )
289
+ coupling = invN * (torch.sin(A) * zc + torch.cos(A) * zs)
290
+ to_center = center_pull * torch.sin(ctr - theta + alpha_i) # [B,T,N]
291
+ theta = theta + cfg.osc_dt * (omega + coupling + to_center + inject) # ring drive
292
+
293
+ feat = torch.cat([torch.cos(theta), torch.sin(theta)], dim=-1) # [B,T,2N]
294
+ out = self.readout(feat) * torch.tanh(self.gate)
295
+ out = self.drop(soft_clamp(out, self.cfg.osc_bound)) # bounded, live grad
296
+
297
+ rec = _viz.get_rec()
298
+ if rec is not None and rec.enabled: # live-viz capture
299
+ th = theta[0, -1] # last-token phases [N]
300
+ cnt = self.onehot.sum(0).clamp(min=1.0)
301
+ zc = (torch.cos(th) @ self.onehot) / cnt # per-ring order param
302
+ zs = (torch.sin(th) @ self.onehot) / cnt
303
+ R = torch.sqrt(zc ** 2 + zs ** 2 + 1e-9)
304
+ psi = torch.atan2(zs, zc)
305
+ rec.log_ring(R.tolist(), psi.tolist(), th.tolist())
306
+ rec.log_quaz_norm(out[0, -1].norm().item())
307
+ rec.flush_spec() # group this layer's routes
308
+ return out
309
+
310
+
311
+ class RotaryEmbedding(nn.Module):
312
+ """RoPE for the rope partition of Q/K (qk_rope_head_dim dims). Ported from v2."""
313
+
314
+ def __init__(self, dim, max_positions=4096, theta=10000.0):
315
+ super().__init__()
316
+ inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim))
317
+ t = torch.arange(max_positions).float()
318
+ freqs = torch.outer(t, inv_freq)
319
+ self.register_buffer("cos_cache", freqs.cos(), persistent=False)
320
+ self.register_buffer("sin_cache", freqs.sin(), persistent=False)
321
+
322
+ def forward(self, x, position_ids):
323
+ cos = self.cos_cache[position_ids].unsqueeze(1)
324
+ sin = self.sin_cache[position_ids].unsqueeze(1)
325
+ d = cos.shape[-1]
326
+ x1, x2 = x[..., :d], x[..., d:]
327
+ return torch.cat([x1 * cos - x2 * sin, x1 * sin + x2 * cos], dim=-1)
328
+
329
+
330
+ class MLADerfXSAAttention(nn.Module):
331
+ """Family attention, ported from model_v2.MLADerfXSAAttention: MLA low-rank
332
+ q/o projections, partial RoPE (nope+rope split), QK-Norm, optional DERF erf
333
+ attention, optional XSA value-subspace removal, GQA. Supports an optional KV
334
+ cache for incremental decoding (cache holds post-RoPE, pre-GQA-expansion k/v
335
+ of shape [B, num_kv_heads, S, head_dim])."""
336
+
337
+ def __init__(self, cfg: QuazimotoConfig):
338
+ super().__init__()
339
+ self.num_heads = cfg.n_head
340
+ self.num_kv_heads = cfg.num_key_value_heads
341
+ self.head_dim = cfg.head_dim
342
+ self.nope_head_dim = cfg.nope_head_dim
343
+ self.use_qk_norm = cfg.use_qk_norm
344
+ self.use_derf = cfg.use_derf
345
+ self.use_xsa = cfg.use_xsa
346
+ self.dropout_p = cfg.dropout
347
+ self.kv_groups = self.num_heads // self.num_kv_heads
348
+ assert self.nope_head_dim + cfg.qk_rope_head_dim == self.head_dim, \
349
+ "nope_head_dim + qk_rope_head_dim must equal head_dim"
350
+
351
+ self.q_a_proj = nn.Linear(cfg.d_model, cfg.q_lora_rank, bias=False)
352
+ self.q_a_norm = RMSNorm(cfg.q_lora_rank, cfg.rms_norm_eps)
353
+ self.q_b_proj = nn.Linear(cfg.q_lora_rank, self.num_heads * self.head_dim, bias=False)
354
+ self.k_proj = nn.Linear(cfg.d_model, self.num_kv_heads * self.head_dim, bias=False)
355
+ self.v_proj = nn.Linear(cfg.d_model, self.num_kv_heads * self.head_dim, bias=False)
356
+ self.o_a_proj = nn.Linear(self.num_heads * self.head_dim, cfg.o_lora_rank, bias=False)
357
+ self.o_b_proj = nn.Linear(cfg.o_lora_rank, cfg.d_model, bias=False)
358
+
359
+ if self.use_qk_norm:
360
+ self.q_norm = RMSNorm(self.head_dim, cfg.rms_norm_eps)
361
+ self.k_norm = RMSNorm(self.head_dim, cfg.rms_norm_eps)
362
+ self.rope = RotaryEmbedding(cfg.qk_rope_head_dim, cfg.max_position_embeddings,
363
+ cfg.rope_theta)
364
+ if self.use_derf:
365
+ self.derf_alpha = nn.Parameter(torch.ones(self.num_heads))
366
+ self.derf_bias = nn.Parameter(torch.zeros(self.num_heads))
367
+ self.derf_gamma = nn.Parameter(torch.ones(self.num_heads))
368
+ for m in (self.q_a_proj, self.q_b_proj, self.k_proj, self.v_proj,
369
+ self.o_a_proj, self.o_b_proj):
370
+ nn.init.normal_(m.weight, std=cfg.initializer_range)
371
+
372
+ def forward(self, x, position_ids, past_kv=None, use_cache=False):
373
+ B, S, _ = x.shape
374
+ q = self.q_b_proj(self.q_a_norm(self.q_a_proj(x)))
375
+ q = q.view(B, S, self.num_heads, self.head_dim).transpose(1, 2)
376
+ k = self.k_proj(x).view(B, S, self.num_kv_heads, self.head_dim).transpose(1, 2)
377
+ v = self.v_proj(x).view(B, S, self.num_kv_heads, self.head_dim).transpose(1, 2)
378
+ if self.use_qk_norm:
379
+ q, k = self.q_norm(q), self.k_norm(k)
380
+
381
+ d = self.nope_head_dim
382
+ q = torch.cat([q[..., :d], self.rope(q[..., d:], position_ids)], dim=-1)
383
+ k = torch.cat([k[..., :d], self.rope(k[..., d:], position_ids)], dim=-1)
384
+
385
+ # KV cache: prepend previously-seen post-RoPE k/v (pre-GQA-expansion so
386
+ # the cache stays GQA-compact), then this step's tokens become the tail.
387
+ if past_kv is not None:
388
+ past_k, past_v = past_kv
389
+ k = torch.cat([past_k, k], dim=2)
390
+ v = torch.cat([past_v, v], dim=2)
391
+ present = (k, v) if use_cache else None
392
+ L = k.size(2) # total keys (cache + current)
393
+
394
+ if self.kv_groups > 1:
395
+ k = k.unsqueeze(2).expand(-1, -1, self.kv_groups, -1, -1).reshape(
396
+ B, self.num_heads, L, self.head_dim)
397
+ v = v.unsqueeze(2).expand(-1, -1, self.kv_groups, -1, -1).reshape(
398
+ B, self.num_heads, L, self.head_dim)
399
+
400
+ # Causal mask over the [S queries x L keys] block. When there is no cache
401
+ # and S == L this is the plain lower triangle; with a cache the S new
402
+ # queries sit at absolute positions [L-S, L) and may attend all keys <=
403
+ # their own position. allowed[i,j] = j <= (L - S + i).
404
+ offset = L - S
405
+ if self.use_derf:
406
+ qpos = torch.arange(S, device=x.device).view(S, 1) + offset
407
+ kpos = torch.arange(L, device=x.device).view(1, L)
408
+ is_masked = (kpos > qpos).unsqueeze(0).unsqueeze(0) # [1,1,S,L]
409
+ scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim)
410
+ safe = scores.masked_fill(is_masked, -10000.0)
411
+ a, b, g = (self.derf_alpha.view(1, -1, 1, 1), self.derf_bias.view(1, -1, 1, 1),
412
+ self.derf_gamma.view(1, -1, 1, 1))
413
+ w = g * torch.erf(a * safe + b)
414
+ w = (w + g) / 2.0
415
+ w = w.masked_fill(is_masked, 0.0)
416
+ w = w / (w.sum(-1, keepdim=True) + 1e-8)
417
+ if self.dropout_p > 0 and self.training:
418
+ w = F.dropout(w, p=self.dropout_p)
419
+ y = torch.matmul(w, v)
420
+ else:
421
+ if offset == 0:
422
+ attn_mask, causal = None, True
423
+ else:
424
+ qpos = torch.arange(S, device=x.device).view(S, 1) + offset
425
+ kpos = torch.arange(L, device=x.device).view(1, L)
426
+ attn_mask = (kpos <= qpos).unsqueeze(0).unsqueeze(0) # bool: True=keep
427
+ causal = False
428
+ y = F.scaled_dot_product_attention(
429
+ q.contiguous(), k.contiguous(), v.contiguous(), attn_mask=attn_mask,
430
+ is_causal=causal, dropout_p=self.dropout_p if self.training else 0.0)
431
+
432
+ rec = _viz.get_rec()
433
+ if rec is not None and rec.enabled: # last-query attention, mean over heads
434
+ sc = (q[:, :, -1:] @ k.transpose(-2, -1)) / math.sqrt(self.head_dim) # [B,H,1,L]
435
+ w = torch.softmax(sc, dim=-1).mean(1)[0, 0] # [L]
436
+ rec.log_attn(getattr(self, "layer_idx", 0), w.tolist())
437
+
438
+ if self.use_xsa:
439
+ # remove each query's component along its OWN value direction; with a
440
+ # cache the queries are the last S of the L cached value positions.
441
+ vq = v[:, :, -S:, :]
442
+ vn = vq / (vq.norm(dim=-1, keepdim=True) + 1e-8)
443
+ y = y - (y * vn).sum(-1, keepdim=True) * vn
444
+
445
+ y = y.transpose(1, 2).contiguous().view(B, S, self.num_heads * self.head_dim)
446
+ out = self.o_b_proj(self.o_a_proj(y))
447
+ return (out, present) if use_cache else out
448
+
449
+
450
+ class HyperConnectionLayer(nn.Module):
451
+ """Softmax pre-mix / post-distribute over hc_mult residual streams (v2)."""
452
+
453
+ def __init__(self, hc_mult):
454
+ super().__init__()
455
+ self.pre_weight = nn.Parameter(torch.linspace(0.5, -0.5, hc_mult) / max(hc_mult, 1))
456
+ self.post_weight = nn.Parameter(torch.linspace(-0.5, 0.5, hc_mult) / max(hc_mult, 1))
457
+
458
+ def pre_op(self, copies):
459
+ w = F.softmax(self.pre_weight, dim=0)
460
+ return (copies * w.view(1, -1, 1, 1)).sum(dim=1)
461
+
462
+ def post_op(self, copies, delta):
463
+ w = F.softmax(self.post_weight, dim=0)
464
+ return copies + delta.unsqueeze(1) * w.view(1, -1, 1, 1)
465
+
466
+
467
+ class HCOutputMix(nn.Module):
468
+ """Learned softmax mix over hc_mult streams at the output (init == mean)."""
469
+
470
+ def __init__(self, hc_mult):
471
+ super().__init__()
472
+ self.weight = nn.Parameter(torch.zeros(hc_mult))
473
+
474
+ def forward(self, copies):
475
+ w = F.softmax(self.weight, dim=0)
476
+ return (copies * w.view(1, -1, 1, 1)).sum(dim=1)
477
+
478
+
479
+ class Layer(nn.Module):
480
+ def __init__(self, cfg):
481
+ super().__init__()
482
+ self.cfg = cfg
483
+ self.use_hc = cfg.use_hyper_connections
484
+ self.use_value_embed = cfg.use_value_embed
485
+ self.attn_norm = RMSNorm(cfg.d_model, cfg.rms_norm_eps)
486
+ self.attn = MLADerfXSAAttention(cfg)
487
+ self.quaz = MycelBlock(cfg) # the Neighbour-Sensing growth sub-layer
488
+ if self.use_hc:
489
+ self.hc_attn = HyperConnectionLayer(cfg.hc_mult)
490
+ self.hc_ffn = HyperConnectionLayer(cfg.hc_mult)
491
+ if self.use_value_embed:
492
+ self.ve_gate = nn.Parameter(torch.zeros(1)) # zero-init -> no-op
493
+
494
+ def forward(self, x, position_ids, token_embed=None, ring_ctl=None,
495
+ past_kv=None, use_cache=False, phase_seed=None):
496
+ h = self.hc_attn.pre_op(x) if self.use_hc else x
497
+ if self.use_value_embed and token_embed is not None:
498
+ h = h + torch.tanh(self.ve_gate) * token_embed
499
+ attn_out = self.attn(self.attn_norm(h), position_ids,
500
+ past_kv=past_kv, use_cache=use_cache)
501
+ present = None
502
+ if use_cache:
503
+ attn_out, present = attn_out
504
+ if self.use_hc:
505
+ x = self.hc_attn.post_op(x, attn_out)
506
+ h = self.hc_ffn.pre_op(x)
507
+ else:
508
+ h = h + attn_out
509
+ ffn_out = self.quaz(h, ring_ctl, phase_seed) # QuazimotoBlock norms internally
510
+ if self.use_hc:
511
+ x = self.hc_ffn.post_op(x, ffn_out)
512
+ else:
513
+ x = h + ffn_out
514
+ return (x, present) if use_cache else x
515
+
516
+
517
+ class QuazimotoLM(nn.Module):
518
+ def __init__(self, cfg: QuazimotoConfig):
519
+ super().__init__()
520
+ self.cfg = cfg
521
+ self.tok = nn.Embedding(cfg.vocab_size, cfg.d_model)
522
+ self.drop = nn.Dropout(cfg.dropout) # positions come from RoPE (no abs embed)
523
+ # frozen Mandelbrot orbit-angle phase table [vocab, n_osc]; non-persistent
524
+ # (deterministic, regenerated on load -> not stored in checkpoints).
525
+ if cfg.use_fractal_phase_seed:
526
+ from fractal import mandelbrot_phase_table, load_phase_table
527
+ pt, mode = load_phase_table(cfg.vocab_size, cfg.n_osc,
528
+ os.path.join(_PKG_DIR, "fractal_phase.pt"))
529
+ if pt is None:
530
+ pt = mandelbrot_phase_table(cfg.vocab_size, cfg.n_osc)
531
+ print(" [fractal] flat Halton seed (run build_fractal_table.py for hierarchical)")
532
+ else:
533
+ print(f" [fractal] loaded {mode} phase table")
534
+ self.register_buffer("fractal_phase", pt, persistent=False)
535
+ self.layers = nn.ModuleList([Layer(cfg) for _ in range(cfg.n_layer)])
536
+ for i, layer in enumerate(self.layers):
537
+ layer.attn.layer_idx = i # for live-viz attention capture
538
+ self.norm_f = RMSNorm(cfg.d_model, cfg.rms_norm_eps)
539
+ self.head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)
540
+ self.head.weight = self.tok.weight # weight tying
541
+ self.logit_scale = nn.Parameter(torch.tensor(1.0)) # family trait: learned logit temp
542
+ self.use_value_embed = cfg.use_value_embed
543
+ self.hc_out_mix = HCOutputMix(cfg.hc_mult) if cfg.use_hyper_connections else None
544
+ self.apply(self._init)
545
+ # engram rings manage their own inits (frozen compressor, std-0.01 tables,
546
+ # DERF bias -4); re-apply them after the global init clobbers those.
547
+ for m in self.modules():
548
+ if isinstance(m, (EngramRing, RingSpecialists, MycelStations)):
549
+ m.family_reinit()
550
+
551
+ # opt-in trait modules -- instantiated AFTER self.apply so their internal
552
+ # zero-inits (MoE down-proj, MTP proj) and gate inits survive the global init.
553
+ # HRM/MoE refine the trunk after the stack (the family's "top membrane
554
+ # refinement" point); MTP/JEPA are train-time aux heads on the trunk rep.
555
+ self.hrm = (HRMRefinementBlock(cfg.d_model, cfg.hrm_dim, cfg.hrm_steps,
556
+ gate_init_open=cfg.hrm_gate_init)
557
+ if cfg.use_hrm else None)
558
+ self.moe = MoESwiGLU(cfg.d_model, cfg.moe_intermediate, cfg.moe_n_routed,
559
+ cfg.moe_n_shared, cfg.moe_top_k) if cfg.use_moe else None
560
+ self.mtp_heads = (nn.ModuleList([MTPHead(cfg.d_model) for _ in range(cfg.mtp_layers)])
561
+ if cfg.use_mtp else None)
562
+ self.jepa = (JEPAPredictorBlock(cfg.d_model, cfg.jepa_pred_dim, cfg.jepa_horizon)
563
+ if cfg.use_jepa else None)
564
+ self.ring_bank = (RingControllerBank(len(cfg.ring_sizes), 4, cfg.ring_ctrl_feat,
565
+ cfg.ring_ctrl_local_lr)
566
+ if cfg.use_ring_controllers else None)
567
+
568
+ # exact kwargs to rebuild on resume (family `family_config` convention)
569
+ self.family_config = {k: getattr(cfg, k) for k in (
570
+ "vocab_size", "n_layer", "n_head", "d_model", "block_size",
571
+ "ring_sizes", "osc_steps", "osc_dt", "readout_mult", "osc_bound", "gate_init_open",
572
+ "use_hrm", "hrm_steps", "hrm_dim", "hrm_gate_init",
573
+ "use_moe", "moe_intermediate", "moe_n_routed",
574
+ "moe_n_shared", "moe_top_k", "use_mtp", "mtp_layers", "mtp_loss_weight",
575
+ "use_jepa", "jepa_pred_dim", "jepa_horizon", "jepa_loss_weight",
576
+ "head_dim", "qk_rope_head_dim", "nope_head_dim", "num_key_value_heads",
577
+ "q_lora_rank", "o_lora_rank", "use_qk_norm", "use_derf", "use_xsa",
578
+ "rope_theta", "max_position_embeddings", "rms_norm_eps", "initializer_range",
579
+ "zloss_coef", "use_value_embed", "use_hyper_connections", "hc_mult",
580
+ "use_rings", "ring_attn_heads", "ring_attn_head_dim", "ring_engram_compress",
581
+ "ring_engram_heads", "ring_engram_table", "ring_engram_ngram",
582
+ "use_ring_controllers", "ring_ctrl_feat", "ring_ctrl_local_lr",
583
+ "use_ring_specialists", "ring_n_specialists", "ring_spec_key_dim",
584
+ "ring_spec_slot_dim", "ring_spec_top_k", "ring_spec_write_lr",
585
+ "use_fractal_phase_seed")}
586
+ n = sum(p.numel() for p in self.parameters())
587
+ print(f"Mycel-LM: {n/1e6:.1f}M params | {cfg.n_tips} tips/token growing in a "
588
+ f"{cfg.mycel_pos_dim}D bounded colony ({cfg.growth_steps} growth steps) | "
589
+ f"traits: {self.active_traits()}")
590
+
591
+ def active_traits(self):
592
+ t = []
593
+ if self.cfg.use_hrm: t.append("HRM")
594
+ if self.cfg.use_moe: t.append("MoE")
595
+ if self.cfg.use_mtp: t.append("MTP")
596
+ if self.cfg.use_jepa: t.append("JEPA")
597
+ if self.cfg.use_rings: t.append("Rings")
598
+ if self.cfg.use_ring_controllers: t.append("Controllers")
599
+ if self.cfg.use_ring_specialists:
600
+ t.append(f"Specialists({self.cfg.ring_n_specialists}/ring)")
601
+ if self.cfg.use_fractal_phase_seed: t.append("FractalSeed")
602
+ return ",".join(t) or "none"
603
+
604
+ def reset_ring_memory(self):
605
+ """Zero every RingSpecialists store. Call between independent prompts so
606
+ the test-time context memory does not carry over from a previous sequence."""
607
+ for m in self.modules():
608
+ if isinstance(m, RingSpecialists):
609
+ m.reset_memory()
610
+
611
+ def set_ring_memory_writing(self, enabled: bool):
612
+ """Toggle online writes to the specialist stores (e.g. freeze the memory
613
+ during evaluation, or disable it to ablate the test-time-write behavior)."""
614
+ for m in self.modules():
615
+ if isinstance(m, RingSpecialists):
616
+ m.write_enabled = enabled
617
+
618
+ def _init(self, m):
619
+ if isinstance(m, nn.Linear):
620
+ nn.init.normal_(m.weight, std=0.02)
621
+ if m.bias is not None:
622
+ nn.init.zeros_(m.bias)
623
+ elif isinstance(m, nn.Embedding):
624
+ nn.init.normal_(m.weight, std=0.02)
625
+
626
+ def _compute_trunk(self, idx, past_key_values=None, use_cache=False):
627
+ """Shared backbone: tokens -> layers -> trunk refinements -> RMSNorm trunk.
628
+ Returns (trunk, presents). Used by both forward() and forward_drafts()."""
629
+ B, T = idx.shape
630
+ # with a cache, the new tokens sit at absolute positions [past_len, past_len+T)
631
+ past_len = past_key_values[0][0].size(2) if past_key_values else 0
632
+ position_ids = (torch.arange(past_len, past_len + T, device=idx.device)
633
+ .unsqueeze(0).expand(B, -1))
634
+ x = self.drop(self.tok(idx))
635
+ token_embed = x if self.use_value_embed else None
636
+ phase_seed = self.fractal_phase[idx] if self.cfg.use_fractal_phase_seed else None
637
+ if self.cfg.use_hyper_connections:
638
+ x = x.unsqueeze(1).expand(-1, self.cfg.hc_mult, -1, -1).clone() # [B,M,T,H]
639
+ presents = [] if use_cache else None
640
+ for i, layer in enumerate(self.layers):
641
+ past = past_key_values[i] if past_key_values is not None else None
642
+ x = layer(x, position_ids, token_embed, self.ring_bank,
643
+ past_kv=past, use_cache=use_cache, phase_seed=phase_seed)
644
+ if use_cache:
645
+ x, present = x
646
+ presents.append(present)
647
+ if self.hc_out_mix is not None:
648
+ x = self.hc_out_mix(x) # learned mix over streams (init == mean)
649
+ # trunk refinements (no-op at init): HRM iterative gated refine, MoE expert mix
650
+ rec = _viz.get_rec()
651
+ if self.hrm is not None:
652
+ hx = self.hrm(x)
653
+ if rec is not None and rec.enabled:
654
+ rec.log_trait("hrm", (hx - x)[0, -1].norm().item())
655
+ x = hx
656
+ if self.moe is not None:
657
+ mx = self.moe(x)
658
+ if rec is not None and rec.enabled:
659
+ rec.log_trait("moe", mx[0, -1].norm().item())
660
+ x = x + mx
661
+ return self.norm_f(x), presents # v2: RMSNorm trunk (no tanh)
662
+
663
+ @torch.no_grad()
664
+ def forward_drafts(self, idx, past_key_values=None, use_cache=False):
665
+ """Speculative-decoding support (DeepSpec-style draft+verify, self-drafted):
666
+ return (main_logits, mtp_logits_list[, presents]) over ALL T positions.
667
+ * main_logits[:, j] predicts the token at position j+1 (verifier)
668
+ * mtp_logits_list[k][:, j] predicts the token at position j+2+k (k=0..mtp-1)
669
+ So one forward yields, at the last position, the genuine next token (main)
670
+ plus `mtp_layers` drafted future tokens to be verified next round."""
671
+ trunk, presents = self._compute_trunk(idx, past_key_values, use_cache)
672
+ main_logits = self.head(trunk) * self.logit_scale
673
+ mtp_logits = ([self.head(h(trunk)) * self.logit_scale for h in self.mtp_heads]
674
+ if self.mtp_heads is not None else [])
675
+ return (main_logits, mtp_logits, presents) if use_cache else (main_logits, mtp_logits)
676
+
677
+ def forward(self, idx, targets=None, past_key_values=None, use_cache=False):
678
+ B, T = idx.shape
679
+ trunk, presents = self._compute_trunk(idx, past_key_values, use_cache)
680
+ logits = self.head(trunk) * self.logit_scale
681
+
682
+ loss, aux = None, {}
683
+ if targets is not None:
684
+ flat = logits.view(-1, logits.size(-1))
685
+ tflat = targets.reshape(-1)
686
+ loss = F.cross_entropy(flat, tflat, ignore_index=-1)
687
+ # z-loss (v2): penalise log^2 of the partition function -> no logit drift
688
+ if self.cfg.zloss_coef > 0:
689
+ valid = tflat != -1
690
+ if valid.any():
691
+ log_z = torch.logsumexp(flat[valid].float(), dim=-1)
692
+ aux["zloss"] = self.cfg.zloss_coef * (log_z ** 2).mean()
693
+ if self.moe is not None and self.moe.last_aux_loss is not None:
694
+ aux["moe_aux"] = self.moe.last_aux_loss
695
+ if self.mtp_heads is not None:
696
+ mtp_total, n_active = 0.0, 0
697
+ for k, head in enumerate(self.mtp_heads, start=1):
698
+ if T - k <= 0:
699
+ break
700
+ mlog = self.head(head(trunk[:, :T - k])) * self.logit_scale
701
+ mtp_total = mtp_total + F.cross_entropy(
702
+ mlog.reshape(-1, self.cfg.vocab_size),
703
+ targets[:, k:].reshape(-1), ignore_index=-1)
704
+ n_active += 1
705
+ if n_active:
706
+ aux["mtp_loss"] = self.cfg.mtp_loss_weight * mtp_total / n_active
707
+ if self.jepa is not None and T > 1:
708
+ jepa_total, n_j = 0.0, 0
709
+ for k in range(1, self.cfg.jepa_horizon + 1):
710
+ if T - k <= 0:
711
+ break
712
+ pred = self.jepa(trunk[:, :T - k], k)
713
+ tgt = trunk[:, k:].detach() # JEPA stop-grad target
714
+ cos = F.cosine_similarity(pred.float(), tgt.float(), dim=-1)
715
+ jepa_total = jepa_total + (1.0 - cos).mean()
716
+ n_j += 1
717
+ if n_j:
718
+ aux["jepa_loss"] = self.cfg.jepa_loss_weight * jepa_total / n_j
719
+ # use_cache returns the per-layer (k, v) presents as a 4th item; without
720
+ # it the signature is the original 3-tuple so train.py is unaffected.
721
+ if use_cache:
722
+ return logits, loss, aux, presents
723
+ return logits, loss, aux
724
+
725
+ @torch.no_grad()
726
+ def generate(self, idx, n_new, temperature=1.0, top_k=None, use_cache=True):
727
+ """Autoregressive sampling with the family's NaN/inf sanitisation: long
728
+ accumulation can push a logit to +/-inf or NaN, whose softmax -> NaN ->
729
+ multinomial fires a CUDA device-side assert. Clamp first, greedy-fallback.
730
+
731
+ With use_cache=True (default) the attention KV cache is kept across steps
732
+ so each new token costs one forward over a single position instead of a
733
+ full recompute. The cache holds absolute positions, so once it fills past
734
+ block_size we drop the cache and fall back to a windowed recompute (RoPE
735
+ positions would otherwise exceed max_position_embeddings)."""
736
+ self.eval()
737
+ past = None
738
+ for _ in range(n_new):
739
+ if use_cache and past is not None and past[0][0].size(2) < self.cfg.block_size:
740
+ step_in = idx[:, -1:] # only the newest token
741
+ logits, _, _, past = self(step_in, past_key_values=past, use_cache=True)
742
+ elif use_cache:
743
+ cond = idx[:, -self.cfg.block_size:] # prefill / re-prime cache
744
+ logits, _, _, past = self(cond, use_cache=True)
745
+ else:
746
+ cond = idx[:, -self.cfg.block_size:]
747
+ logits, _, _ = self(cond)
748
+ lg = torch.nan_to_num(logits[:, -1, :].float(), nan=0.0,
749
+ posinf=1e4, neginf=-1e4) / max(temperature, 1e-6)
750
+ if top_k:
751
+ k = min(top_k, lg.size(-1))
752
+ v, _ = torch.topk(lg, k)
753
+ lg[lg < v[:, [-1]]] = -float("inf")
754
+ probs = torch.softmax(lg, dim=-1)
755
+ if not torch.isfinite(probs).all() or float(probs.sum()) <= 0.0:
756
+ nxt = torch.argmax(lg, dim=-1, keepdim=True)
757
+ else:
758
+ nxt = torch.multinomial(probs, 1)
759
+ idx = torch.cat([idx, nxt], dim=1)
760
+ return idx
761
+
762
+
763
+ if __name__ == "__main__":
764
+ # exercise the base model and all trait modules at once
765
+ cfg = QuazimotoConfig(use_hrm=True, use_moe=True, use_mtp=True, use_jepa=True)
766
+ model = QuazimotoLM(cfg)
767
+ x = torch.randint(0, cfg.vocab_size, (2, 64))
768
+ logits, loss, aux = model(x, x)
769
+ total = loss + sum(aux.values())
770
+ print("forward ok:", tuple(logits.shape), "loss", round(loss.item(), 3),
771
+ "| aux:", {k: round(float(v), 4) for k, v in aux.items()})
772
+ total.backward()
773
+ print("backward ok")
mycel.py ADDED
@@ -0,0 +1,178 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ mycel.py -- the Neighbour-Sensing (fungal colony growth) channel-mixing block for
3
+ Mycel-LM, the structural analog of Quazimoto's Kuramoto oscillator block.
4
+
5
+ Meskauskas/Fricker/Moore (2004) Neighbour-Sensing model: each hyphal TIP is an agent
6
+ with a position and a growth vector; tips generate a density FIELD; each tip SENSES
7
+ the local field and steers its growth (negative autotropism -- grow away from the
8
+ colony's own density) plus persistence. We turn that into a differentiable mixer:
9
+
10
+ hidden -> N tips (position p in a BOUNDED latent region + growth direction v)
11
+ -> a few differentiable growth steps: sense the density field, steer, move,
12
+ RE-CLAMP into the bounded region (the colony can't grow unbounded)
13
+ -> readout of [p, v, sensed-density] -> hidden, behind a family gate.
14
+
15
+ Traits are "spread through the mycelium" (MycelStations): tiny memory specialists sit
16
+ at fixed anchor positions in the bounded region; a tip interacts with a station by
17
+ PROXIMITY, which is emergent from where it grew -- so which tips use which trait
18
+ "comes to be" during growth rather than being assigned to a fixed index.
19
+
20
+ The density field is computed against F learnable field centres (a low-rank sample of
21
+ the field) so cost is O(N*F) per step, not O(N^2) -- the same trick that keeps the
22
+ oscillator block's mean-field cheap.
23
+ """
24
+ import math
25
+ import torch
26
+ import torch.nn as nn
27
+ import torch.nn.functional as F
28
+
29
+ from family import soft_clamp, RMSNorm
30
+ import instrument as _viz
31
+
32
+
33
+ class MycelStations(nn.Module):
34
+ """Memory specialists placed at anchor positions in the colony's bounded region.
35
+ A tip reads/writes a station by proximity (softmax over -distance). Slow weights
36
+ (anchors, encoders, decoder) are learned; the stores are test-time fast memory
37
+ (buffers). Zero-init scale => no-op at start (family contract)."""
38
+
39
+ def __init__(self, n_stations, pos_dim, d_model, slot_dim=64, top_k=2,
40
+ write_lr=0.1, bound=10.0):
41
+ super().__init__()
42
+ self.n = n_stations
43
+ self.top_k = min(top_k, n_stations)
44
+ self.write_lr, self.bound, self.write_enabled = write_lr, bound, True
45
+ self.anchors = nn.Parameter(torch.randn(n_stations, pos_dim) * 0.5) # station sites
46
+ self.log_bw = nn.Parameter(torch.zeros(1)) # proximity bandwidth
47
+ self.in_enc = nn.Linear(d_model, slot_dim, bias=False) # frozen write encoders
48
+ self.val_enc = nn.Linear(d_model, slot_dim, bias=False)
49
+ self.out_dec = nn.Linear(slot_dim, pos_dim, bias=False) # station -> steering drive
50
+ self.scale = nn.Parameter(torch.zeros(1))
51
+ self.register_buffer("store_in", torch.zeros(n_stations, slot_dim))
52
+ self.register_buffer("store_out", torch.zeros(n_stations, slot_dim))
53
+ for m in (self.in_enc, self.val_enc, self.out_dec):
54
+ nn.init.normal_(m.weight, std=0.02)
55
+ self.in_enc.weight.requires_grad_(False)
56
+ self.val_enc.weight.requires_grad_(False)
57
+
58
+ def family_reinit(self):
59
+ nn.init.normal_(self.anchors, std=0.5)
60
+ nn.init.zeros_(self.scale)
61
+ for m in (self.in_enc, self.val_enc):
62
+ nn.init.normal_(m.weight, std=0.02); m.weight.requires_grad_(False)
63
+
64
+ def reset_memory(self):
65
+ self.store_in.zero_(); self.store_out.zero_()
66
+
67
+ def forward(self, p, h):
68
+ # p: [B,T,N,pos_dim] tip positions ; h: [B,T,d] token hidden
69
+ store_in, store_out = self.store_in.clone(), self.store_out.clone() # noqa: F841
70
+ bw = F.softplus(self.log_bw).clamp(min=1e-3)
71
+ # tip->station squared distance via the expansion (no [..,N,n,pos_dim] intermediate)
72
+ p2 = (p * p).sum(-1, keepdim=True) # [B,T,N,1]
73
+ a2 = (self.anchors * self.anchors).sum(-1) # [n]
74
+ d2 = (p2 + a2 - 2.0 * (p @ self.anchors.t())).clamp(min=0.0) # [B,T,N,n]
75
+ route = torch.softmax(-d2 / bw, dim=-1)
76
+ if self.top_k < self.n:
77
+ tv = torch.topk(route, self.top_k, dim=-1).values
78
+ route = route.masked_fill(route < tv[..., [-1]], 0.0)
79
+ route = route / route.sum(-1, keepdim=True).clamp(min=1e-6)
80
+ drive = self.out_dec(route @ store_out) * torch.tanh(self.scale) # [B,T,N,pos_dim]
81
+
82
+ if self.write_enabled and self.write_lr > 0: # stations absorb nearby tips
83
+ with torch.no_grad():
84
+ occ = route.reshape(-1, self.n) # [BTN, n] proximity mass
85
+ info_in = self.in_enc(h).unsqueeze(2).expand(-1, -1, p.size(2), -1).reshape(-1, self.in_enc.out_features)
86
+ info_v = self.val_enc(h).unsqueeze(2).expand(-1, -1, p.size(2), -1).reshape(-1, self.val_enc.out_features)
87
+ denom = occ.sum(0).clamp(min=1e-3).unsqueeze(1)
88
+ a = self.write_lr
89
+ self.store_in.mul_(1 - a).add_(a * (occ.t() @ info_in) / denom).clamp_(-self.bound, self.bound)
90
+ self.store_out.mul_(1 - a).add_(a * (occ.t() @ info_v) / denom).clamp_(-self.bound, self.bound)
91
+ return drive
92
+
93
+
94
+ class MycelBlock(nn.Module):
95
+ """Neighbour-Sensing growth mixer: hidden -> growing tips in a bounded region ->
96
+ readout. Structural analog of QuazimotoBlock.forward(x, ring_ctl, phase_seed)."""
97
+
98
+ def __init__(self, cfg):
99
+ super().__init__()
100
+ self.cfg = cfg
101
+ self.N, self.pd, self.F = cfg.n_tips, cfg.mycel_pos_dim, cfg.mycel_field_centers
102
+ self.norm = RMSNorm(cfg.d_model)
103
+ self.to_pos = nn.Linear(cfg.d_model, self.N * self.pd) # initial tip positions
104
+ self.to_dir = nn.Linear(cfg.d_model, self.N * self.pd) # initial growth vectors
105
+ self.centers = nn.Parameter(torch.randn(self.F, self.pd) * 0.5) # field sample sites
106
+ self.log_persist = nn.Parameter(torch.tensor(1.4)) # sigmoid -> ~0.8 persistence
107
+ self.tropism = nn.Parameter(torch.zeros(1)) # negative-autotropism strength
108
+ self.log_field_bw = nn.Parameter(torch.zeros(1))
109
+ hidden = cfg.readout_mult * cfg.d_model
110
+ self.readout = nn.Sequential(
111
+ nn.Linear(self.N * (2 * self.pd + 1), hidden), nn.GELU(),
112
+ nn.Linear(hidden, cfg.d_model))
113
+ self.drop = nn.Dropout(cfg.dropout)
114
+ go = math.atanh(min(cfg.gate_init_open, 0.9)) if cfg.gate_init_open > 0 else 0.0
115
+ self.gate = nn.Parameter(torch.full((1,), go))
116
+ for m in self.modules():
117
+ if isinstance(m, nn.Linear):
118
+ nn.init.normal_(m.weight, std=0.02)
119
+ if m.bias is not None:
120
+ nn.init.zeros_(m.bias)
121
+ if cfg.use_fractal_phase_seed:
122
+ self.pos_seed_gate = nn.Parameter(torch.zeros(1)) # zero-init -> no-op
123
+ self.use_stations = cfg.use_ring_specialists
124
+ if self.use_stations:
125
+ self.stations = MycelStations(cfg.mycel_n_stations, self.pd, cfg.d_model,
126
+ cfg.ring_spec_slot_dim, cfg.ring_spec_top_k,
127
+ cfg.ring_spec_write_lr, cfg.osc_bound)
128
+
129
+ def _clamp(self, p):
130
+ b = self.cfg.osc_bound
131
+ return b * torch.tanh(p / b) # bound the colony region
132
+
133
+ @staticmethod
134
+ def _cdist2(p, c):
135
+ # squared distances [..,N,F] via ||p||^2 + ||c||^2 - 2 p.c -- NO [..,N,F,pd]
136
+ # intermediate (that broadcast is the memory hog; this keeps it O(N*F)).
137
+ p2 = (p * p).sum(-1, keepdim=True) # [..,N,1]
138
+ c2 = (c * c).sum(-1) # [F]
139
+ return (p2 + c2 - 2.0 * (p @ c.t())).clamp(min=0.0) # [..,N,F]
140
+
141
+ def _grow(self, p, v, h):
142
+ cfg = self.cfg
143
+ pers = torch.sigmoid(self.log_persist)
144
+ trop = torch.tanh(self.tropism)
145
+ bw = F.softplus(self.log_field_bw).clamp(min=1e-3)
146
+ for _ in range(cfg.growth_steps):
147
+ K = torch.exp(-self._cdist2(p, self.centers) / bw) # [B,T,N,F]
148
+ w = K.mean(-2, keepdim=True) * K # rho_f * K_if [B,T,N,F]
149
+ # steer away from dense centres, factored to avoid the [..,N,F,pd] tensor:
150
+ # sum_f w_if (p_i - c_f) = p_i * sum_f w_if - (w @ centres)
151
+ away = p * w.sum(-1, keepdim=True) - torch.matmul(w, self.centers) # [B,T,N,pd]
152
+ v = pers * v + trop * away
153
+ if self.use_stations:
154
+ v = v + self.stations(p, h) # traits steer nearby tips
155
+ p = self._clamp(p + cfg.growth_dt * v)
156
+ dens = torch.exp(-self._cdist2(p, self.centers) / bw).mean(-1, keepdim=True)
157
+ return p, v, dens
158
+
159
+ def forward(self, x, ring_ctl=None, phase_seed=None):
160
+ cfg = self.cfg
161
+ B, T, _ = x.shape
162
+ h = self.norm(x)
163
+ p = self._clamp(self.to_pos(h).view(B, T, self.N, self.pd))
164
+ if phase_seed is not None: # optional fractal position seed
165
+ seed = phase_seed[..., :self.N * self.pd].view(B, T, self.N, self.pd)
166
+ p = self._clamp(p + torch.tanh(self.pos_seed_gate) * seed)
167
+ v = self.to_dir(h).view(B, T, self.N, self.pd)
168
+ p, v, dens = self._grow(p, v, h)
169
+ feat = torch.cat([p, v, dens], dim=-1).flatten(2) # [B,T,N*(2pd+1)]
170
+ out = self.drop(soft_clamp(self.readout(feat) * torch.tanh(self.gate), cfg.osc_bound))
171
+
172
+ rec = _viz.get_rec()
173
+ if rec is not None and rec.enabled: # live-viz: colony spread
174
+ spread = dens[0, -1, :, 0]
175
+ rec.log_ring([float(spread.mean())], [0.0], p[0, -1].flatten().tolist())
176
+ rec.log_quaz_norm(out[0, -1].norm().item())
177
+ rec.flush_spec()
178
+ return out
opd_teacher.py ADDED
@@ -0,0 +1,79 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ opd_teacher.py -- frozen teacher (facebook/MobileLLM-R1-140M-base) for on-policy
3
+ distillation into Quazimoto-LM.
4
+
5
+ The teacher and student use DIFFERENT tokenizers, so a token-level KL is undefined.
6
+ This scorer instead returns the teacher's per-token log-probabilities of a piece of
7
+ text together with each token's BYTE span, so train_opd.py can re-aggregate that
8
+ dense signal onto the student's own tokens (byte-space alignment). The teacher is
9
+ never trained; it only scores the student's rollouts.
10
+
11
+ Requires: transformers, torch. The teacher is ~140M params -> fits on a small GPU
12
+ (or CPU, slowly). Run on the rig that has the weights / network access.
13
+ """
14
+
15
+ import torch
16
+ import torch.nn.functional as F
17
+
18
+
19
+ class Teacher:
20
+ def __init__(self, name="facebook/MobileLLM-R1-140M-base", device="cuda",
21
+ dtype=torch.float32):
22
+ from transformers import AutoModelForCausalLM, AutoTokenizer
23
+ self.device = device
24
+ self.tok = AutoTokenizer.from_pretrained(name, trust_remote_code=True)
25
+ try: # newer transformers: dtype=
26
+ self.model = AutoModelForCausalLM.from_pretrained(
27
+ name, dtype=dtype, trust_remote_code=True)
28
+ except TypeError: # older transformers: torch_dtype=
29
+ self.model = AutoModelForCausalLM.from_pretrained(
30
+ name, torch_dtype=dtype, trust_remote_code=True)
31
+ self.model = self.model.to(device).eval()
32
+ self.has_offsets = bool(getattr(self.tok, "is_fast", False))
33
+ if not self.has_offsets:
34
+ print(" [teacher] WARNING: slow tokenizer (no offset mapping) -> "
35
+ "byte alignment unavailable; train_opd will fall back to seq reward.")
36
+
37
+ @torch.no_grad()
38
+ def score(self, text):
39
+ """Return (byte_logprob, n_bytes): a per-BYTE teacher log-prob vector for
40
+ `text` (length = len(text.encode('utf-8'))). Each teacher token's log-prob
41
+ (of being generated given its prefix) is spread uniformly over the bytes it
42
+ covers. byte_logprob[k] is the teacher's per-byte score for byte k -- the
43
+ dense signal the student is distilled toward.
44
+
45
+ If offsets are unavailable, returns (mean_logprob_scalar, n_bytes) with the
46
+ scalar broadcast (sequence-level fallback)."""
47
+ raw = text.encode("utf-8")
48
+ nb = len(raw)
49
+ if nb == 0:
50
+ return torch.zeros(0), 0
51
+ enc = self.tok(text, return_offsets_mapping=self.has_offsets,
52
+ return_tensors="pt", add_special_tokens=False)
53
+ ids = enc["input_ids"].to(self.device)
54
+ T = ids.size(1)
55
+ if T < 2:
56
+ return torch.zeros(nb), nb
57
+ logits = self.model(ids).logits[0].float() # [T, V]
58
+ logp = F.log_softmax(logits, dim=-1)
59
+ # logprob of token t (t>=1) given prefix = logp[t-1, ids[t]]
60
+ tok_lp = torch.zeros(T)
61
+ tok_lp[1:] = logp[torch.arange(T - 1), ids[0, 1:]].cpu()
62
+
63
+ if not self.has_offsets:
64
+ return torch.full((nb,), float(tok_lp[1:].mean())), nb
65
+
66
+ # char offsets -> byte offsets (UTF-8 widths of the prefix)
67
+ char_to_byte = [0]
68
+ for ch in text:
69
+ char_to_byte.append(char_to_byte[-1] + len(ch.encode("utf-8")))
70
+ offs = enc["offset_mapping"][0].tolist() # [(c0,c1), ...]
71
+ byte_lp = torch.zeros(nb)
72
+ for t in range(T):
73
+ c0, c1 = offs[t]
74
+ if c1 <= c0:
75
+ continue
76
+ b0, b1 = char_to_byte[c0], char_to_byte[c1]
77
+ span = max(b1 - b0, 1)
78
+ byte_lp[b0:b1] = tok_lp[t] / span # spread token logprob over its bytes
79
+ return byte_lp, nb
requirements.txt ADDED
@@ -0,0 +1,11 @@
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Core (inference, generation, visualizer, health-check)
2
+ torch>=2.1
3
+ numpy
4
+ transformers>=4.40 # SpikeTokenizer subclasses PreTrainedTokenizer
5
+
6
+ # Training / fine-tuning (train.py, train_sft.py, distillation)
7
+ datasets>=2.18 # streamed pretrain / SFT blends
8
+ huggingface_hub>=0.23 # gated dataset + repo access
9
+
10
+ # Optional
11
+ # safetensors # not required; checkpoints are torch .pt
special_tokens.py ADDED
@@ -0,0 +1,85 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ special_tokens.py -- central, *append-only* registry of special tokens for the
3
+ SpikeWhale length-max tokenizer.
4
+
5
+ WHY THIS FILE EXISTS
6
+ --------------------
7
+ The base vocab (tokenizer.json) is 16384 contiguous ids:
8
+ 0..3 -> <pad> <unk> <bos> <eos>
9
+ 4..259 -> the 256 raw bytes
10
+ 260.. -> learned byte-merges
11
+ Adding tokens "without breaking the model" has exactly one rule:
12
+
13
+ ***APPEND ONLY. NEVER REORDER OR REMOVE AN EXISTING ID.***
14
+
15
+ Every existing id keeps pointing at the same embedding row and the same logit
16
+ column, so the model's behaviour on already-seen tokens is bit-for-bit
17
+ unchanged. New tokens are appended at ids >= 16384 and their embedding / lm_head
18
+ / mtp rows are freshly initialised (near-zero contribution) so they are
19
+ no-ops until you train them.
20
+
21
+ To stay tensor-core friendly the final vocab is padded up to a multiple of
22
+ `VOCAB_MULTIPLE` (128) with `<|reserved_N|>` slots. Those reserves let you name
23
+ *future* tokens later by editing the registry WITHOUT another model resize, as
24
+ long as the total stays <= the padded size.
25
+
26
+ HOW TO ADD MORE LATER
27
+ ---------------------
28
+ Append new names to NAMED_SPECIAL_TOKENS (at the END), then either:
29
+ * if you still have <|reserved_*|> slots free, just rename a reserved id in
30
+ tokenizer.json (no model change needed), or
31
+ * re-run add_special_tokens.py to grow + re-pad the vocab (model resized).
32
+ """
33
+
34
+ # Tensor-core / matmul friendly vocab alignment. 16384 is already 128*128.
35
+ VOCAB_MULTIPLE = 128
36
+
37
+ # ---------------------------------------------------------------------------
38
+ # The universal named set. ORDER IS PERMANENT -- append only, never reorder.
39
+ # Mixing the common conventions so the same model can do chat, reasoning,
40
+ # agentic tool use, and code infilling.
41
+ # ---------------------------------------------------------------------------
42
+ NAMED_SPECIAL_TOKENS = [
43
+ # ChatML turn framing
44
+ "<|im_start|>",
45
+ "<|im_end|>",
46
+ # Reasoning / scratchpad
47
+ "<think>",
48
+ "</think>",
49
+ # Explicit solution block
50
+ "<begin_solution>",
51
+ "<end_solution>",
52
+ # Agentic tool calling
53
+ "<tool_call>",
54
+ "</tool_call>",
55
+ "<tool_response>",
56
+ "</tool_response>",
57
+ # Role markers (usable standalone or inside an im_start header)
58
+ "<|system|>",
59
+ "<|user|>",
60
+ "<|assistant|>",
61
+ # Fill-in-the-middle (code)
62
+ "<|fim_prefix|>",
63
+ "<|fim_middle|>",
64
+ "<|fim_suffix|>",
65
+ # Generic document separator
66
+ "<|endoftext|>",
67
+ ]
68
+
69
+
70
+ def build_special_token_list(base_vocab_size: int,
71
+ multiple: int = VOCAB_MULTIPLE):
72
+ """
73
+ Return the ordered list of tokens to APPEND after `base_vocab_size`:
74
+ the named set followed by enough <|reserved_N|> slots to pad the final
75
+ vocab size up to the next multiple of `multiple`.
76
+
77
+ The returned list's element i gets id (base_vocab_size + i).
78
+ """
79
+ tokens = list(NAMED_SPECIAL_TOKENS)
80
+ target = base_vocab_size + len(tokens)
81
+ # round up to the next multiple (or stay put if already aligned)
82
+ padded = ((target + multiple - 1) // multiple) * multiple
83
+ n_reserved = padded - target
84
+ tokens += [f"<|reserved_{i}|>" for i in range(n_reserved)]
85
+ return tokens
spike_tokenizer.py ADDED
@@ -0,0 +1,124 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ spike_tokenizer.py -- HuggingFace-compatible wrapper for the custom
3
+ byte-level "length-max" (greedy longest-match) tokenizer in tokenizer.json.
4
+
5
+ The raw tokenizer.json is NOT a HuggingFace `tokenizers` file; it is a plain
6
+ dict {vocab, vocab_size, max_token_len, algorithm:"length-max"}. This wrapper
7
+ makes it loadable by AutoTokenizer.from_pretrained / save_pretrained and
8
+ exposes encode/decode + the bos/eos/pad/unk ids the training scripts expect.
9
+
10
+ Encoding scheme (verified): byte-level. Text is UTF-8 encoded, each byte mapped
11
+ to its latin-1 character, then greedily matched against the vocab using the
12
+ longest key that matches at each position (max key length = max_token_len).
13
+ """
14
+ import json, os
15
+ from typing import List, Optional
16
+
17
+ # We only use transformers' tokenizer base class. Stop it from probing
18
+ # TensorFlow/Flax (whose builds are often incompatible with numpy 2 / protobuf and
19
+ # would crash the import). setdefault so an explicit user override still wins.
20
+ os.environ.setdefault("USE_TF", "0")
21
+ os.environ.setdefault("USE_FLAX", "0")
22
+
23
+ from transformers import PreTrainedTokenizer
24
+
25
+
26
+ class SpikeTokenizer(PreTrainedTokenizer):
27
+ vocab_files_names = {"vocab_file": "tokenizer.json"}
28
+ model_input_names = ["input_ids"]
29
+
30
+ def __init__(self, vocab_file=None, **kwargs):
31
+ with open(vocab_file, "r", encoding="utf-8") as f:
32
+ data = json.load(f)
33
+ self._vocab = data["vocab"] # str -> id
34
+ self._ids_to_tokens = {i: t for t, i in self._vocab.items()}
35
+ self.max_token_len = int(data.get("max_token_len", 24))
36
+ # length-bucketed keys for fast greedy match (longest length first)
37
+ self._lengths = sorted({len(k) for k in self._vocab}, reverse=True)
38
+
39
+ # Appended special tokens (im_start / <think> / <begin_solution> / ...).
40
+ # They already live in self._vocab at their real ids; we hand them to the
41
+ # HF base class as `additional_special_tokens` so its AddedToken trie:
42
+ # (1) splits them out ATOMICALLY before our byte-level greedy match
43
+ # (verified: each maps back to its existing vocab id, no phantom id), and
44
+ # (2) drops them on decode(skip_special_tokens=True).
45
+ # The set is stored in tokenizer.json under "special_tokens" so it
46
+ # survives save_pretrained/from_pretrained round-trips.
47
+ self._extra_specials = [
48
+ t for t in data.get("special_tokens", []) if t in self._vocab
49
+ ]
50
+ if self._extra_specials:
51
+ existing = list(kwargs.get("additional_special_tokens", []) or [])
52
+ merged = existing + [t for t in self._extra_specials if t not in existing]
53
+ kwargs["additional_special_tokens"] = merged
54
+
55
+ kwargs.setdefault("bos_token", "<bos>")
56
+ kwargs.setdefault("eos_token", "<eos>")
57
+ kwargs.setdefault("unk_token", "<unk>")
58
+ kwargs.setdefault("pad_token", "<pad>")
59
+ super().__init__(**kwargs)
60
+
61
+ @property
62
+ def vocab_size(self) -> int:
63
+ return len(self._vocab)
64
+
65
+ def get_vocab(self):
66
+ return dict(self._vocab)
67
+
68
+ # --- core byte-level greedy tokenization ---
69
+ def _tokenize(self, text: str) -> List[str]:
70
+ s = text.encode("utf-8").decode("latin-1") # one char per byte
71
+ out, i, n = [], 0, len(s)
72
+ while i < n:
73
+ matched = None
74
+ hi = min(self.max_token_len, n - i)
75
+ for L in range(hi, 0, -1):
76
+ sub = s[i:i + L]
77
+ if sub in self._vocab:
78
+ matched = sub
79
+ break
80
+ if matched is None: # single byte always exists in vocab
81
+ matched = s[i]
82
+ out.append(matched)
83
+ i += len(matched)
84
+ return out
85
+
86
+ def _convert_token_to_id(self, token: str) -> int:
87
+ return self._vocab.get(token, self._vocab["<unk>"])
88
+
89
+ def _convert_id_to_token(self, index: int) -> str:
90
+ return self._ids_to_tokens.get(index, "<unk>")
91
+
92
+ def convert_tokens_to_string(self, tokens: List[str]) -> str:
93
+ # transformers 5.x hands the FULL token list here (special tokens
94
+ # included; skip_special_tokens is already applied upstream via
95
+ # convert_ids_to_tokens). So we can't just byte-decode everything: a
96
+ # special token like "<|im_start|>" is a literal marker, not latin-1
97
+ # bytes. Decode runs of ordinary byte-tokens together (needed so
98
+ # multi-byte UTF-8 sequences reassemble) and emit any special token
99
+ # inline as its literal string.
100
+ specials = {"<pad>", "<unk>", "<bos>", "<eos>", *self._extra_specials}
101
+ out, buf = [], []
102
+ for tok in tokens:
103
+ if tok in specials:
104
+ if buf:
105
+ out.append("".join(buf).encode("latin-1").decode("utf-8", errors="replace"))
106
+ buf = []
107
+ out.append(tok)
108
+ else:
109
+ buf.append(tok)
110
+ if buf:
111
+ out.append("".join(buf).encode("latin-1").decode("utf-8", errors="replace"))
112
+ return "".join(out)
113
+
114
+ def save_vocabulary(self, save_directory: str, filename_prefix: Optional[str] = None):
115
+ os.makedirs(save_directory, exist_ok=True)
116
+ fn = (filename_prefix + "-" if filename_prefix else "") + "tokenizer.json"
117
+ path = os.path.join(save_directory, fn)
118
+ with open(path, "w", encoding="utf-8") as f:
119
+ json.dump({"vocab": self._vocab, "vocab_size": self.vocab_size,
120
+ "max_token_len": self.max_token_len,
121
+ "algorithm": "length-max",
122
+ "special_tokens": list(self._extra_specials)},
123
+ f, ensure_ascii=False)
124
+ return (path,)
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
train.bat ADDED
@@ -0,0 +1,24 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ @echo off
2
+ REM Mycel-LM pretraining -- the Neighbour-Sensing (fungal colony growth) mixer, same
3
+ REM harness/traits/data as the Quazimoto lm3 run, so the two are a controlled swap
4
+ REM (mixer is the only architectural difference). Ring-specific traits (--use-rings,
5
+ REM --use-ring-controllers) are OMITTED: the mycelium has no rings -- its analog is
6
+ REM --use-ring-specialists, which places the trait STATIONS through the colony.
7
+ cd /d "%~dp0"
8
+ REM the growth loop is activation-heavy; expandable_segments cuts fragmentation.
9
+ set PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True
10
+ python train.py ^
11
+ --steps 130000 ^
12
+ --batch 4 ^
13
+ --block 512 ^
14
+ --math-frac 0.25 ^
15
+ --use-hrm ^
16
+ --use-moe ^
17
+ --use-mtp ^
18
+ --use-jepa ^
19
+ --use-ring-specialists ^
20
+ --use-fractal-phase-seed ^
21
+ --amp ^
22
+ --stream ^
23
+ --resume
24
+ pause
train.py ADDED
@@ -0,0 +1,297 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Train Quazimoto-LM with the SpikeWhale tokenizer and the family's AdamW regime.
3
+
4
+ Tokenizer: the custom SpikeWhale length-max byte tokenizer (tokenizer.json +
5
+ spike_tokenizer.py) bundled in this package -- self-contained, no external folder.
6
+ Data is tokenized fresh in Python from a local UTF-8 .txt (or a synthetic fallback).
7
+
8
+ Optimizer matches the family's logged-stable choice (LEVEL6_PLAN): AdamW, peak lr
9
+ 3e-4, linear warmup, cosine decay to min_lr_frac*peak (0.3 floor), grad-clip 1.0.
10
+ MuonEq-R is deliberately NOT used here -- it caused weight runaway in this family.
11
+ """
12
+
13
+ import argparse, time, math, os, sys
14
+ import numpy as np
15
+ import torch
16
+ from model import QuazimotoLM, QuazimotoConfig
17
+
18
+ # SpikeWhale tokenizer is bundled in this package (self-contained, no level6 needed).
19
+ PKG_DIR = os.path.dirname(os.path.abspath(__file__))
20
+
21
+
22
+ def load_tokenizer(tok_dir):
23
+ sys.path.insert(0, tok_dir)
24
+ from spike_tokenizer import SpikeTokenizer
25
+ tok = SpikeTokenizer(vocab_file=os.path.join(tok_dir, "tokenizer.json"))
26
+ return tok
27
+
28
+
29
+ def load_ids(path, tok):
30
+ if path and os.path.exists(path):
31
+ with open(path, "r", encoding="utf-8", errors="replace") as f:
32
+ text = f.read()
33
+ else:
34
+ text = ("the quazimoto oscillator learns to synchronize phases across rings. "
35
+ "coupled clocks fall into step when coupling is strong enough. ") * 4000
36
+ ids = np.array(tok.encode(text, add_special_tokens=False), dtype=np.int64)
37
+ n = int(0.9 * len(ids))
38
+ return ids[:n], ids[n:]
39
+
40
+
41
+ def get_batch(ids, bs, block, device):
42
+ ix = np.random.randint(0, len(ids) - block - 1, size=bs)
43
+ x = np.stack([ids[i:i + block] for i in ix])
44
+ y = np.stack([ids[i + 1:i + 1 + block] for i in ix])
45
+ return (torch.from_numpy(x).to(device), torch.from_numpy(y).to(device))
46
+
47
+
48
+ def _doc_text(ex):
49
+ """Pull the text field from a dataset example, tolerating naming differences
50
+ across datasets (text / content / raw_content / document), else the first
51
+ non-trivial string value."""
52
+ for k in ("text", "content", "raw_content", "document"):
53
+ v = ex.get(k)
54
+ if isinstance(v, str) and v.strip():
55
+ return v
56
+ for v in ex.values():
57
+ if isinstance(v, str) and len(v.strip()) > 1:
58
+ return v
59
+ return None
60
+
61
+
62
+ def _stream_docs(path, config, split):
63
+ """Infinite generator of non-empty doc strings from a streaming HF dataset.
64
+ Re-opens the stream when a shard pass is exhausted so training never runs dry.
65
+ (Mirrors the level6 train_snn_stream._stream_docs convention.)"""
66
+ from datasets import load_dataset
67
+ while True:
68
+ ds = load_dataset(path, name=config, split=split, streaming=True)
69
+ any_yielded = False
70
+ for ex in ds:
71
+ text = _doc_text(ex)
72
+ if text:
73
+ any_yielded = True
74
+ yield text
75
+ if not any_yielded:
76
+ raise RuntimeError(f"stream {path}:{config} yielded no usable text")
77
+
78
+
79
+ def _blend_sources(args):
80
+ """Return the active (label, path, config, weight) sources for the blend.
81
+ Default mix: 35% Ultra-FineWeb-L3 / 25% FineWeb-Edu / 25% FineMath / 15% PretrainNew."""
82
+ srcs = [
83
+ ("ultra", args.ultra_dataset, args.ultra_config, args.ultra_frac),
84
+ ("edu", args.fineweb_dataset, args.fineweb_config, args.edu_frac),
85
+ ("math", args.math_dataset, args.math_config, args.math_frac),
86
+ ("pretrain", args.pretrain_dataset, args.pretrain_config, args.pretrain_frac),
87
+ ]
88
+ return [s for s in srcs if s[3] > 0]
89
+
90
+
91
+ def doc_generator(args):
92
+ """Blend multiple HF datasets by per-document probability (weights normalized)."""
93
+ import random
94
+ rng = random.Random(args.seed)
95
+ srcs = _blend_sources(args)
96
+ gens = [_stream_docs(p, c, args.split) for (_, p, c, _) in srcs]
97
+ weights = [s[3] for s in srcs]
98
+ tot = sum(weights)
99
+ cum, acc = [], 0.0
100
+ for w in weights: # cumulative thresholds for sampling
101
+ acc += w / tot
102
+ cum.append(acc)
103
+ while True:
104
+ r = rng.random()
105
+ gi = next(k for k, c in enumerate(cum) if r <= c)
106
+ yield next(gens[gi]) + "\n"
107
+
108
+
109
+ def stream_batches(tok, args, device):
110
+ """Infinite [B,T] next-token batches from the streamed blend via a token buffer.
111
+ Docs are EOS-separated and packed; each batch is B contiguous T+1 windows."""
112
+ eos = tok.eos_token_id
113
+ sep = [eos] if eos is not None else []
114
+ need = args.batch * (args.block + 1)
115
+ buf, docs = [], doc_generator(args)
116
+ srcs = _blend_sources(args)
117
+ tot = sum(s[3] for s in srcs)
118
+ blend = " / ".join(f"{int(round(100*s[3]/tot))}% {s[1]}" for s in srcs)
119
+ print(f"streaming blend: {blend}")
120
+ while True:
121
+ while len(buf) < need:
122
+ buf.extend(tok.encode(next(docs), add_special_tokens=False))
123
+ buf.extend(sep)
124
+ chunk = np.array(buf[:need], dtype=np.int64).reshape(args.batch, args.block + 1)
125
+ del buf[:need]
126
+ x = torch.from_numpy(chunk[:, :-1]).to(device)
127
+ y = torch.from_numpy(chunk[:, 1:]).to(device)
128
+ yield x, y
129
+
130
+
131
+ def save_ckpt(model, tok_vocab, step, path, opt=None):
132
+ if not path:
133
+ return
134
+ os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
135
+ ckpt = {"model": model.state_dict(), "family_config": model.family_config,
136
+ "vocab_size": tok_vocab, "step": step}
137
+ if opt is not None:
138
+ ckpt["optim"] = opt.state_dict() # so --resume continues the optimizer too
139
+ torch.save(ckpt, path)
140
+ print(f" saved checkpoint -> {path} (step {step})")
141
+
142
+
143
+ def find_latest_ckpt(out_path):
144
+ """Auto-locate a checkpoint to resume from: prefer the exact --out file, else
145
+ the most recently modified *.pt in the same folder. Returns None if none."""
146
+ if out_path and os.path.isfile(out_path):
147
+ return out_path
148
+ folder = os.path.dirname(out_path) or "."
149
+ if not os.path.isdir(folder):
150
+ return None
151
+ pts = [os.path.join(folder, f) for f in os.listdir(folder) if f.endswith(".pt")]
152
+ return max(pts, key=os.path.getmtime) if pts else None
153
+
154
+
155
+ def lr_at(step, peak, warmup, total, min_frac):
156
+ if step < warmup:
157
+ return peak * step / max(warmup, 1)
158
+ if step >= total:
159
+ return peak * min_frac
160
+ prog = (step - warmup) / max(total - warmup, 1)
161
+ return peak * (min_frac + (1 - min_frac) * 0.5 * (1 + math.cos(math.pi * prog)))
162
+
163
+
164
+ def main():
165
+ p = argparse.ArgumentParser()
166
+ p.add_argument("--data", default="")
167
+ p.add_argument("--tok-dir", default=PKG_DIR, help="dir with bundled tokenizer.json + spike_tokenizer.py")
168
+ p.add_argument("--steps", type=int, default=200)
169
+ p.add_argument("--batch", type=int, default=8)
170
+ p.add_argument("--block", type=int, default=256)
171
+ p.add_argument("--lr", type=float, default=3e-4) # family peak
172
+ p.add_argument("--warmup", type=int, default=1000)
173
+ p.add_argument("--min-lr-frac", type=float, default=0.3) # family floor (~9e-5)
174
+ p.add_argument("--weight-decay", type=float, default=0.01)
175
+ p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
176
+ p.add_argument("--use-hrm", action="store_true")
177
+ p.add_argument("--use-moe", action="store_true")
178
+ p.add_argument("--use-mtp", action="store_true")
179
+ p.add_argument("--mtp-layers", type=int, default=4, help="MTP draft-head depth (spec-decode)")
180
+ p.add_argument("--use-jepa", action="store_true")
181
+ p.add_argument("--use-rings", action="store_true")
182
+ p.add_argument("--use-ring-controllers", action="store_true")
183
+ p.add_argument("--use-ring-specialists", action="store_true",
184
+ help="per-ring MoE memory specialists with test-time input/output stores")
185
+ p.add_argument("--use-fractal-phase-seed", action="store_true",
186
+ help="seed oscillator phases from each token's Mandelbrot orbit (gated)")
187
+ p.add_argument("--out", default=os.path.join(PKG_DIR, "chkpt", "quazimoto.pt")) # checkpoint file path
188
+ p.add_argument("--ckpt-every", type=int, default=250)
189
+ p.add_argument("--resume", action="store_true",
190
+ help="auto-find the latest checkpoint in the --out folder and continue training")
191
+ # streaming blend (HF datasets); off by default so local runs stay dependency-free.
192
+ # Default 4-way mix: 35% Ultra-FineWeb-L3 / 25% FineWeb-Edu / 25% FineMath / 15% PretrainNew.
193
+ p.add_argument("--stream", action="store_true", help="stream the weighted dataset blend")
194
+ p.add_argument("--ultra-frac", type=float, default=0.35)
195
+ p.add_argument("--ultra-dataset", default="openbmb/Ultra-FineWeb-L3")
196
+ p.add_argument("--ultra-config", default="Ultra-FineWeb-L3-en-Multi-Style-Synthetic")
197
+ p.add_argument("--edu-frac", type=float, default=0.25)
198
+ p.add_argument("--fineweb-dataset", default="HuggingFaceFW/fineweb-edu")
199
+ p.add_argument("--fineweb-config", default="sample-10BT")
200
+ p.add_argument("--math-frac", type=float, default=0.25)
201
+ p.add_argument("--math-dataset", default="HuggingFaceTB/finemath")
202
+ p.add_argument("--math-config", default="finemath-4plus")
203
+ p.add_argument("--pretrain-frac", type=float, default=0.15)
204
+ p.add_argument("--pretrain-dataset", default="nvidia/Nemotron-Pretraining-Specialized-v1.1")
205
+ p.add_argument("--pretrain-config", default="Nemotron-Pretraining-Formal-Logic")
206
+ p.add_argument("--split", default="train")
207
+ p.add_argument("--seed", type=int, default=0)
208
+ p.add_argument("--amp", action="store_true",
209
+ help="bf16 autocast (biggest single-GPU speedup; also frees memory)")
210
+ args = p.parse_args()
211
+
212
+ # single-GPU speed: TF32 matmuls (free, ~fp32 accuracy on Ampere+/Blackwell).
213
+ if args.device == "cuda":
214
+ torch.backends.cuda.matmul.allow_tf32 = True
215
+ torch.backends.cudnn.allow_tf32 = True
216
+
217
+ tok = load_tokenizer(args.tok_dir)
218
+ cfg = QuazimotoConfig(vocab_size=tok.vocab_size, block_size=args.block,
219
+ use_hrm=args.use_hrm, use_moe=args.use_moe,
220
+ use_mtp=args.use_mtp, mtp_layers=args.mtp_layers,
221
+ use_jepa=args.use_jepa,
222
+ use_rings=args.use_rings,
223
+ use_ring_controllers=args.use_ring_controllers,
224
+ use_ring_specialists=args.use_ring_specialists,
225
+ use_fractal_phase_seed=args.use_fractal_phase_seed)
226
+ model = QuazimotoLM(cfg).to(args.device)
227
+ print(f"tokenizer: SpikeWhale (vocab {tok.vocab_size})")
228
+
229
+ if args.stream:
230
+ batches = stream_batches(tok, args, args.device) # infinite blended stream
231
+ train = None
232
+ else:
233
+ train, _ = load_ids(args.data, tok)
234
+ print(f"data: {len(train)} local tokens | device {args.device}")
235
+
236
+ opt = torch.optim.AdamW(model.parameters(), lr=args.lr, betas=(0.9, 0.95),
237
+ weight_decay=args.weight_decay,
238
+ fused=(args.device == "cuda")) # fused CUDA optimizer step
239
+
240
+ start_step = 1
241
+ if args.resume:
242
+ ckpt_path = find_latest_ckpt(args.out)
243
+ if ckpt_path is None:
244
+ print(f"--resume: no checkpoint found near {args.out}; starting fresh.")
245
+ else:
246
+ ck = torch.load(ckpt_path, map_location=args.device, weights_only=False)
247
+ miss, unexp = model.load_state_dict(ck["model"], strict=False)
248
+ if miss: print(f" [resume warn] missing keys: {len(miss)} (e.g. {miss[:2]})")
249
+ if unexp: print(f" [resume warn] unexpected keys: {len(unexp)} (e.g. {unexp[:2]})")
250
+ if "optim" in ck:
251
+ try:
252
+ opt.load_state_dict(ck["optim"])
253
+ except ValueError as e:
254
+ print(f" [resume warn] optimizer state not restored ({e}); using fresh optimizer.")
255
+ start_step = int(ck.get("step", 0)) + 1
256
+ print(f"resumed from {ckpt_path} at step {ck.get('step')} -> continuing at {start_step}")
257
+ if start_step > args.steps:
258
+ print(f" already at/past --steps ({args.steps}); nothing to do.")
259
+
260
+ model.train()
261
+ t0 = time.time()
262
+ for step in range(start_step, args.steps + 1):
263
+ lr = lr_at(step, args.lr, args.warmup, args.steps, args.min_lr_frac)
264
+ for g in opt.param_groups:
265
+ g["lr"] = lr
266
+ x, y = (next(batches) if args.stream
267
+ else get_batch(train, args.batch, args.block, args.device))
268
+ # bf16 autocast: ~1.5-2x on Blackwell tensor cores. bf16 keeps fp32 range so
269
+ # no grad scaler is needed, and it halves activation memory (room for bigger
270
+ # --batch). Params stay fp32; the oscillator trig/erf math is safe in bf16.
271
+ with torch.autocast("cuda", dtype=torch.bfloat16,
272
+ enabled=args.amp and args.device == "cuda"):
273
+ _, loss, aux = model(x, y)
274
+ total = loss + sum(aux.values()) # add active trait aux losses
275
+ opt.zero_grad(set_to_none=True)
276
+ total.backward()
277
+ gn = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
278
+ if not torch.isfinite(gn): # family grad-finiteness guard
279
+ print(f"step {step}: non-finite grad norm, skipping")
280
+ opt.zero_grad(set_to_none=True)
281
+ continue
282
+ opt.step()
283
+ if step % 10 == 0 or step == 1:
284
+ dt = time.time() - t0
285
+ extra = " ".join(f"{k} {float(v):.3f}" for k, v in aux.items())
286
+ bank = getattr(model, "ring_bank", None)
287
+ surp = f" surprise {float(bank.last_surprise):.3f}" if bank is not None else ""
288
+ print(f"step {step:4d} | loss {loss.item():.3f} | "
289
+ f"bpt {loss.item()/math.log(2):.3f} | lr {lr:.2e} | {extra}{surp} | {dt:.1f}s")
290
+ if args.ckpt_every and step % args.ckpt_every == 0:
291
+ save_ckpt(model, tok.vocab_size, step, args.out, opt)
292
+ save_ckpt(model, tok.vocab_size, args.steps, args.out, opt)
293
+ print("done.")
294
+
295
+
296
+ if __name__ == "__main__":
297
+ main()
train_opd.py ADDED
@@ -0,0 +1,290 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ train_opd.py -- On-Policy Distillation of MobileLLM-R1-140M-base (teacher) into
3
+ Quazimoto-LM (student), across different tokenizers.
4
+
5
+ Loop (per step):
6
+ 1. draw P text prompts from the streaming blend; take each prompt's first
7
+ --prompt-len student tokens.
8
+ 2. the STUDENT samples G on-policy rollouts per prompt (its own tokens).
9
+ 3. decode each rollout to text; the frozen TEACHER scores that text and returns a
10
+ per-BYTE log-prob; we re-aggregate it onto the student's tokens (byte-space
11
+ alignment in opd_teacher.Teacher.score) -> a DENSE per-student-token reward.
12
+ 4. GRPO-normalize rewards within each prompt's group of G rollouts -> advantages.
13
+ 5. policy-gradient update WITH a KL-to-reference anchor:
14
+ loss = -mean(advantage_t * logprob_t) + kl_beta * KL(student || frozen_ref)
15
+ The KL term is mandatory, not optional: unanchored on-policy RL against a reward
16
+ model reliably collapses into degenerate high-reward text (reward hacking) --
17
+ reward shoots past the teacher's own self-entropy while generations turn to mush.
18
+
19
+ WATCH: teacher-reward/byte should rise then PLATEAU while kl stays bounded (small,
20
+ non-exploding). If reward keeps climbing past the teacher's self-score AND kl blows
21
+ up, lower --lr or raise --kl-beta. Do a ~300-step canary and eyeball a generation
22
+ before committing to a long run.
23
+
24
+ This keeps OPD's dense per-token teacher signal on the student's own samples while
25
+ remaining tokenizer-agnostic (a token-level KL is impossible across the two vocabs).
26
+ Resumable; checkpoints are interchangeable with train.py's format.
27
+
28
+ Run on a rig with the teacher weights + transformers installed:
29
+ python train_opd.py --steps 5000 --prompts 4 --group 4 --gen-len 64 \
30
+ --student-ckpt chkpt/quazimoto.pt --device cuda --stream
31
+ """
32
+
33
+ import argparse, math, os, time
34
+ import numpy as np
35
+ import torch
36
+ import torch.nn.functional as F
37
+
38
+ from model import QuazimotoLM, QuazimotoConfig
39
+ import train as T
40
+
41
+ PKG_DIR = os.path.dirname(os.path.abspath(__file__))
42
+
43
+
44
+ # --------------------------------------------------------------------------- #
45
+ # student helpers
46
+ # --------------------------------------------------------------------------- #
47
+ def token_byte_lens(tok):
48
+ """Map student id -> number of bytes its decoding contributes (for alignment)."""
49
+ specials = set(getattr(tok, "_extra_specials", []))
50
+ ids2 = tok._ids_to_tokens
51
+ cache = {}
52
+ for tid, s in ids2.items():
53
+ if tid < 4 or s in specials:
54
+ cache[tid] = len(s.encode("utf-8"))
55
+ else:
56
+ cache[tid] = len(s.encode("latin-1"))
57
+ return cache
58
+
59
+
60
+ @torch.no_grad()
61
+ def sample_rollouts(model, cfg, idx, gen_len, temperature, top_k):
62
+ """Sample gen_len continuation tokens for a BATCH of equal-length starts.
63
+ idx: [B, Pl] on device. Returns [B, gen_len]. Uses the batched KV cache: the
64
+ starts are prefilled once, then each new token is a single-position forward.
65
+ Batching all prompts x rollouts here collapses what were P separate generation
66
+ loops into ONE, which is the main OPD throughput lever."""
67
+ past, gen = None, []
68
+ for _ in range(gen_len):
69
+ if past is not None and past[0][0].size(2) < cfg.block_size:
70
+ logits, _, _, past = model(idx[:, -1:], past_key_values=past, use_cache=True)
71
+ else: # prefill / re-prime
72
+ logits, _, _, past = model(idx[:, -cfg.block_size:], use_cache=True)
73
+ lg = torch.nan_to_num(logits[:, -1, :].float(), nan=0.0, posinf=1e4, neginf=-1e4)
74
+ lg = lg / max(temperature, 1e-6)
75
+ if top_k:
76
+ v = torch.topk(lg, min(top_k, lg.size(-1)), dim=-1).values
77
+ lg = lg.masked_fill(lg < v[:, [-1]], float("-inf"))
78
+ nxt = torch.multinomial(torch.softmax(lg, dim=-1), 1)
79
+ idx = torch.cat([idx, nxt], dim=1)
80
+ gen.append(nxt)
81
+ return torch.cat(gen, dim=1) # [B, gen_len]
82
+
83
+
84
+ def student_logprobs(model, cfg, full, prompt_len):
85
+ """Per-token log-probs (WITH grad) the student assigns to the rollout tokens,
86
+ batched. full: [B, Pl+gen_len] (prompt ++ rollout). Returns [B, gen_len].
87
+ Assumes Pl+gen_len <= block_size (the normal OPD regime)."""
88
+ B, L = full.shape
89
+ logits, _, _ = model(full[:, -cfg.block_size:])
90
+ logp = F.log_softmax(logits.float(), dim=-1) # [B, L, V]
91
+ gl = L - prompt_len
92
+ tokens = full[:, prompt_len:] # [B, gen_len]
93
+ pred = logp[:, prompt_len - 1:L - 1, :] # dist predicting each rollout token
94
+ return pred.gather(-1, tokens.unsqueeze(-1)).squeeze(-1) # [B, gen_len]
95
+
96
+
97
+ def rollout_token_byte_spans(blen, rollout):
98
+ """Byte span [a,b) each rollout token occupies in the decoded text -- O(T) via
99
+ the precomputed per-id byte lengths (byte tokens concatenate, so cumulative
100
+ lengths match the teacher's text bytes; the caller guards the tail with b<=nb
101
+ in case a trailing partial multi-byte char shifts the decoded length)."""
102
+ spans, prev = [], 0
103
+ for tid in rollout.tolist():
104
+ nb = prev + blen.get(int(tid), 0)
105
+ spans.append((prev, nb))
106
+ prev = nb
107
+ return spans
108
+
109
+
110
+ # --------------------------------------------------------------------------- #
111
+ def main():
112
+ p = argparse.ArgumentParser()
113
+ p.add_argument("--tok-dir", default=PKG_DIR)
114
+ p.add_argument("--student-ckpt", default=os.path.join(PKG_DIR, "chkpt", "quazimoto.pt"))
115
+ p.add_argument("--teacher", default="facebook/MobileLLM-R1-140M-base")
116
+ p.add_argument("--steps", type=int, default=2000)
117
+ p.add_argument("--prompts", type=int, default=4, help="distinct prompts per step (P)")
118
+ p.add_argument("--group", type=int, default=4, help="rollouts per prompt (G, GRPO group)")
119
+ p.add_argument("--prompt-len", type=int, default=32)
120
+ p.add_argument("--gen-len", type=int, default=64)
121
+ p.add_argument("--temperature", type=float, default=1.0)
122
+ p.add_argument("--top-k", type=int, default=50)
123
+ p.add_argument("--lr", type=float, default=1e-5)
124
+ p.add_argument("--warmup", type=int, default=200)
125
+ p.add_argument("--min-lr-frac", type=float, default=0.3)
126
+ p.add_argument("--weight-decay", type=float, default=0.0)
127
+ p.add_argument("--kl-beta", type=float, default=0.1,
128
+ help="KL-to-reference anchor weight (prevents reward hacking; keep > 0)")
129
+ p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
130
+ p.add_argument("--teacher-device", default="", help="default: same as --device")
131
+ p.add_argument("--out", default=os.path.join(PKG_DIR, "chkpt", "quazimoto_opd.pt"))
132
+ p.add_argument("--ckpt-every", type=int, default=200)
133
+ p.add_argument("--resume", action="store_true")
134
+ # streaming blend (prompts source) -- same flags/defaults as train.py
135
+ p.add_argument("--stream", action="store_true")
136
+ p.add_argument("--data", default="")
137
+ for name, default in [("ultra-frac", 0.35), ("edu-frac", 0.25), ("math-frac", 0.25),
138
+ ("pretrain-frac", 0.15)]:
139
+ p.add_argument(f"--{name}", type=float, default=default)
140
+ p.add_argument("--ultra-dataset", default="openbmb/Ultra-FineWeb-L3")
141
+ p.add_argument("--ultra-config", default="Ultra-FineWeb-L3-en-Multi-Style-Synthetic")
142
+ p.add_argument("--fineweb-dataset", default="HuggingFaceFW/fineweb-edu")
143
+ p.add_argument("--fineweb-config", default="sample-10BT")
144
+ p.add_argument("--math-dataset", default="HuggingFaceTB/finemath")
145
+ p.add_argument("--math-config", default="finemath-4plus")
146
+ p.add_argument("--pretrain-dataset", default="Quazim0t0/PretrainNew")
147
+ p.add_argument("--pretrain-config", default=None)
148
+ p.add_argument("--split", default="train")
149
+ p.add_argument("--seed", type=int, default=0)
150
+ args = p.parse_args()
151
+
152
+ dev = args.device
153
+ tdev = args.teacher_device or dev
154
+ tok = T.load_tokenizer(args.tok_dir)
155
+
156
+ # student: rebuild from its checkpoint and continue training
157
+ ck = torch.load(args.student_ckpt, map_location=dev, weights_only=False)
158
+ cfg = QuazimotoConfig(**ck["family_config"])
159
+ student = QuazimotoLM(cfg)
160
+ student.load_state_dict(ck["model"], strict=False)
161
+ student.to(dev)
162
+ print(f"student: step {ck.get('step')} ckpt, vocab {cfg.vocab_size}")
163
+
164
+ # frozen REFERENCE policy = the pristine pre-OPD student. The KL anchor keeps the
165
+ # trained policy from drifting into degenerate high-reward text (reward hacking).
166
+ # Built from ck["model"] BEFORE any resume mutates `student`, so it stays stable.
167
+ reference = QuazimotoLM(cfg)
168
+ reference.load_state_dict(ck["model"], strict=False)
169
+ reference.to(dev).eval()
170
+ for pr in reference.parameters():
171
+ pr.requires_grad_(False)
172
+ # BUGFIX: the specialist memory stores are written on every forward regardless of
173
+ # train/eval mode, so without this the "frozen" reference silently drifts as it
174
+ # scores rollouts -> a moving anchor, weakening the KL. Freeze + reset its memory.
175
+ if hasattr(reference, "set_ring_memory_writing"):
176
+ reference.set_ring_memory_writing(False)
177
+ reference.reset_ring_memory()
178
+ print(f"reference: frozen pre-OPD student (KL anchor, beta={args.kl_beta})")
179
+
180
+ from opd_teacher import Teacher
181
+ teacher = Teacher(args.teacher, device=tdev)
182
+ print(f"teacher: {args.teacher} on {tdev} (frozen)")
183
+
184
+ blen = token_byte_lens(tok) # id -> byte length, for O(T) rollout byte spans
185
+ opt = torch.optim.AdamW(student.parameters(), lr=args.lr, betas=(0.9, 0.95),
186
+ weight_decay=args.weight_decay)
187
+
188
+ # OPD has its OWN step counter (independent of the student's pretraining step).
189
+ # Resume ONLY from this run's own output file -- never from the base checkpoint
190
+ # that may sit in the same folder (that's what --student-ckpt is for).
191
+ start_step = 1
192
+ if args.resume:
193
+ rp = args.out if os.path.isfile(args.out) else None
194
+ if rp:
195
+ r = torch.load(rp, map_location=dev, weights_only=False)
196
+ student.load_state_dict(r["model"], strict=False)
197
+ if "optim" in r:
198
+ try: opt.load_state_dict(r["optim"])
199
+ except ValueError: pass
200
+ start_step = int(r.get("step", 0)) + 1
201
+ print(f"resumed OPD from {rp} at step {r.get('step')} -> {start_step}")
202
+ else:
203
+ print(f"--resume: no prior OPD checkpoint at {args.out}; starting OPD at step 1")
204
+ if start_step > args.steps:
205
+ print(f" nothing to do: OPD already at step {start_step-1} >= --steps {args.steps}. "
206
+ f"Raise --steps to train more."); return
207
+
208
+ # prompt source: first --prompt-len student tokens of streamed docs
209
+ if args.stream:
210
+ docs = T.doc_generator(args)
211
+ else:
212
+ ids_all, _ = T.load_ids(args.data, tok)
213
+ docs = None
214
+
215
+ def next_prompt():
216
+ if docs is not None:
217
+ while True:
218
+ t = tok.encode(next(docs), add_special_tokens=False)
219
+ if len(t) >= args.prompt_len + 1:
220
+ return t[:args.prompt_len]
221
+ i = np.random.randint(0, len(ids_all) - args.prompt_len - 1)
222
+ return ids_all[i:i + args.prompt_len].tolist()
223
+
224
+ student.train()
225
+ t0 = time.time()
226
+ for step in range(start_step, args.steps + 1):
227
+ lr = T.lr_at(step, args.lr, args.warmup, args.steps, args.min_lr_frac)
228
+ for g in opt.param_groups:
229
+ g["lr"] = lr
230
+
231
+ opt.zero_grad(set_to_none=True)
232
+ P, G, GL = args.prompts, args.group, args.gen_len
233
+
234
+ # ONE batched generation over all P*G rollouts (prompts repeated G times)
235
+ prompts = torch.tensor([next_prompt() for _ in range(P)], device=dev) # [P, Pl]
236
+ idx0 = prompts.repeat_interleave(G, dim=0) # [P*G, Pl]
237
+ rolls = sample_rollouts(student, cfg, idx0, GL, args.temperature, args.top_k)
238
+
239
+ # dense per-token teacher reward via byte alignment (one teacher fwd per rollout)
240
+ rewards = torch.zeros(P * G, GL, device=dev)
241
+ for b in range(P * G):
242
+ text = tok.decode(rolls[b].tolist(), skip_special_tokens=False)
243
+ byte_lp, nb = teacher.score(text)
244
+ byte_lp = byte_lp.to(dev)
245
+ for j, (a, e) in enumerate(rollout_token_byte_spans(blen, rolls[b])):
246
+ if e > a and e <= nb:
247
+ rewards[b, j] = byte_lp[a:e].mean()
248
+
249
+ # GRPO: normalize within each prompt's group of G rollouts
250
+ rw = rewards.view(P, G, GL)
251
+ adv = ((rw - rw.mean(dim=(1, 2), keepdim=True))
252
+ / (rw.std(dim=(1, 2), keepdim=True) + 1e-6)).view(P * G, GL)
253
+
254
+ # ONE batched student forward+backward for the policy gradient
255
+ full = torch.cat([idx0, rolls], dim=1) # [P*G, Pl+GL]
256
+ lp = student_logprobs(student, cfg, full, args.prompt_len) # [P*G, GL] (grad)
257
+
258
+ # KL-to-reference anchor: penalize divergence of the trained policy from the
259
+ # frozen pre-OPD student. Without this, on-policy RL against the teacher reward
260
+ # collapses into degenerate high-reward text. k3 estimator (low-variance,
261
+ # non-negative) of KL(student || reference), per token.
262
+ with torch.no_grad():
263
+ ref_lp = student_logprobs(reference, cfg, full, args.prompt_len) # [P*G, GL]
264
+ logr = (ref_lp - lp)
265
+ kl = (torch.exp(logr) - logr - 1.0).mean() # >= 0
266
+
267
+ pg = -(adv.detach() * lp).mean()
268
+ loss = pg + args.kl_beta * kl
269
+ loss.backward()
270
+ step_loss = float(pg.detach()); step_reward = float(rewards.mean()) * P
271
+ step_kl = float(kl.detach())
272
+
273
+ gn = torch.nn.utils.clip_grad_norm_(student.parameters(), 1.0)
274
+ if torch.isfinite(gn):
275
+ opt.step()
276
+ else:
277
+ print(f"step {step}: non-finite grad, skip")
278
+
279
+ if step % 10 == 0 or step == 1:
280
+ dt = time.time() - t0
281
+ print(f"step {step:5d} | pg-loss {step_loss:+.4f} | teacher-reward/byte "
282
+ f"{step_reward/args.prompts:+.3f} | kl {step_kl:.3f} | lr {lr:.2e} | {dt:.1f}s")
283
+ if args.ckpt_every and step % args.ckpt_every == 0:
284
+ T.save_ckpt(student, tok.vocab_size, step, args.out, opt)
285
+ T.save_ckpt(student, tok.vocab_size, args.steps, args.out, opt)
286
+ print("done.")
287
+
288
+
289
+ if __name__ == "__main__":
290
+ main()
train_sft.bat ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ @echo off
2
+ REM Mycel-LM SFT: fine-tune the pretrained mycelium base (chkpt/quazimoto.pt) on the
3
+ REM 4-way chat blend (ultrachat + ultrafeedback + UltraData + OpenThoughts).
4
+ REM --init-ckpt loads the base; --resume continues from chkpt/quazimoto_sft.pt.
5
+ REM If you OOM: drop --batch to 4, or --block to 512.
6
+ cd /d "%~dp0"
7
+ python train_sft.py ^
8
+ --init-ckpt chkpt/quazimoto.pt ^
9
+ --out chkpt/quazimoto_sft.pt ^
10
+ --steps 4000 ^
11
+ --batch 8 ^
12
+ --block 1024 ^
13
+ --lr 2e-5 ^
14
+ --stream ^
15
+ --resume
16
+ pause
train_sft.py ADDED
@@ -0,0 +1,235 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ train_sft.py -- supervised fine-tuning of Quazimoto-LM on UltraChat (chat SFT).
3
+
4
+ Fine-tunes the pretrained base checkpoint on a blended chat-SFT mix (mlabonne
5
+ ultrachat_200k_sft + ultrafeedback-sft + openbmb UltraData-SFT-2605/Knowledge +
6
+ volcanos OpenThoughts2-1M-ShortThink reasoning), normalizing ChatML and ShareGPT
7
+ schemas, rendered in the student's ChatML, with loss MASKED to assistant turns (std
8
+ instruction SFT -- the model learns to produce assistant responses, not to parrot
9
+ the user). Plain next-token cross-entropy -> stable, no reward model, no hacking.
10
+
11
+ Masking detail: with a byte-merge tokenizer you must encode the whole conversation
12
+ ONCE (encoding turn-by-turn would change byte-merges at the boundaries), then derive
13
+ the mask by walking the atomic ChatML special tokens: loss is on for everything
14
+ after an <|assistant|> marker up to and including the next <|im_end|>.
15
+
16
+ Example:
17
+ python train_sft.py --init-ckpt chkpt/quazimoto.pt --steps 4000 --batch 8 \
18
+ --block 1024 --lr 2e-5 --stream --device cuda
19
+ """
20
+ import argparse, math, os, time
21
+ import numpy as np
22
+ import torch
23
+
24
+ from model import QuazimotoLM, QuazimotoConfig
25
+ import train as T
26
+
27
+ PKG_DIR = os.path.dirname(os.path.abspath(__file__))
28
+ ROLE_TOK = {"system": "<|system|>", "user": "<|user|>", "assistant": "<|assistant|>"}
29
+
30
+
31
+ def make_renderer(tok):
32
+ vocab = tok.get_vocab()
33
+ asst_id = vocab.get("<|assistant|>")
34
+ imend_id = vocab.get("<|im_end|>")
35
+ if asst_id is None or imend_id is None:
36
+ raise ValueError("tokenizer missing ChatML special tokens (<|assistant|>/<|im_end|>)")
37
+
38
+ def render(messages, max_len):
39
+ """conversation -> (ids, loss_mask); loss_mask=1 on assistant response tokens."""
40
+ parts = []
41
+ for m in messages:
42
+ role = ROLE_TOK.get(m.get("role", "user"), "<|user|>")
43
+ parts.append(f"<|im_start|>{role}\n{m.get('content','')}<|im_end|>\n")
44
+ ids = tok.encode("".join(parts), add_special_tokens=False)[:max_len]
45
+ mask, in_asst = [], False
46
+ for tid in ids:
47
+ if tid == asst_id: # header marker -> no loss on it
48
+ in_asst = True; mask.append(0); continue
49
+ mask.append(1 if in_asst else 0) # response tokens (incl the closing im_end)
50
+ if tid == imend_id and in_asst:
51
+ in_asst = False
52
+ return ids, mask
53
+ return render
54
+
55
+
56
+ _ROLE_ALIAS = {"human": "user", "user": "user", "gpt": "assistant",
57
+ "assistant": "assistant", "model": "assistant", "bot": "assistant",
58
+ "system": "system"}
59
+
60
+
61
+ def _messages_of(ex):
62
+ """Pull and NORMALIZE a conversation to [{role,content},...] from a dataset row,
63
+ tolerating both ChatML ({role,content}) and ShareGPT ({from,value}) schemas and
64
+ different field names (messages / conversations / conversation / chosen)."""
65
+ for k in ("messages", "conversations", "conversation", "chosen"):
66
+ v = ex.get(k)
67
+ if not (isinstance(v, list) and v and isinstance(v[0], dict)):
68
+ continue
69
+ out = []
70
+ for m in v:
71
+ if "content" in m: # ChatML {role, content}
72
+ role, content = m.get("role", "user"), m.get("content", "")
73
+ elif "value" in m: # ShareGPT {from, value}
74
+ role, content = m.get("from", "user"), m.get("value", "")
75
+ else:
76
+ continue
77
+ out.append({"role": _ROLE_ALIAS.get(str(role).lower(), "user"),
78
+ "content": content or ""})
79
+ if out:
80
+ return out
81
+ return None
82
+
83
+
84
+ def sft_stream(dataset, config, split):
85
+ """Infinite stream of chat conversations (list-of-messages) from one SFT dataset
86
+ (config may be None for single-config datasets)."""
87
+ from datasets import load_dataset
88
+ while True:
89
+ ds = load_dataset(dataset, name=config, split=split, streaming=True)
90
+ any_y = False
91
+ for ex in ds:
92
+ msgs = _messages_of(ex)
93
+ if msgs:
94
+ any_y = True; yield msgs
95
+ if not any_y:
96
+ raise RuntimeError(f"{dataset}:{config}:{split} yielded no conversations")
97
+
98
+
99
+ def blended_convos(sources, seed=0):
100
+ """Blend several SFT datasets by per-conversation probability.
101
+ sources: list of (dataset, config, split, weight)."""
102
+ import random
103
+ rng = random.Random(seed)
104
+ gens = [sft_stream(d, c, s) for (d, c, s, _) in sources]
105
+ weights = [w for (_, _, _, w) in sources]
106
+ tot = sum(weights)
107
+ cum, acc = [], 0.0
108
+ for w in weights:
109
+ acc += w / tot; cum.append(acc)
110
+ print("SFT blend: " + " / ".join(f"{int(round(100*w/tot))}% {d}"
111
+ for (d, _, _, w) in sources))
112
+ while True:
113
+ r = rng.random()
114
+ yield next(gens[next(k for k, c in enumerate(cum) if r <= c)])
115
+
116
+
117
+ def main():
118
+ p = argparse.ArgumentParser()
119
+ p.add_argument("--tok-dir", default=PKG_DIR)
120
+ p.add_argument("--init-ckpt", default=os.path.join(PKG_DIR, "chkpt", "quazimoto.pt"),
121
+ help="pretrained base checkpoint to fine-tune from")
122
+ p.add_argument("--steps", type=int, default=4000)
123
+ p.add_argument("--batch", type=int, default=8)
124
+ p.add_argument("--block", type=int, default=1024)
125
+ p.add_argument("--lr", type=float, default=2e-5, help="SFT lr (lower than pretraining)")
126
+ p.add_argument("--warmup", type=int, default=100)
127
+ p.add_argument("--min-lr-frac", type=float, default=0.1)
128
+ p.add_argument("--weight-decay", type=float, default=0.0)
129
+ p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
130
+ p.add_argument("--out", default=os.path.join(PKG_DIR, "chkpt", "quazimoto_sft.pt"))
131
+ p.add_argument("--ckpt-every", type=int, default=250)
132
+ p.add_argument("--resume", action="store_true")
133
+ p.add_argument("--stream", action="store_true", help="stream the SFT blend")
134
+ # blended SFT sources (default ~equal quarters)
135
+ p.add_argument("--ultrachat-frac", type=float, default=0.25)
136
+ p.add_argument("--ultrachat-dataset", default="mlabonne/ultrachat_200k_sft")
137
+ p.add_argument("--ultrachat-config", default=None)
138
+ p.add_argument("--ultrachat-split", default="train_sft")
139
+ p.add_argument("--feedback-frac", type=float, default=0.25)
140
+ p.add_argument("--feedback-dataset", default="mlabonne/ultrafeedback-sft")
141
+ p.add_argument("--feedback-config", default=None)
142
+ p.add_argument("--feedback-split", default="train")
143
+ p.add_argument("--ultradata-frac", type=float, default=0.25)
144
+ p.add_argument("--ultradata-dataset", default="openbmb/UltraData-SFT-2605")
145
+ p.add_argument("--ultradata-config", default="Knowledge")
146
+ p.add_argument("--ultradata-split", default="no_think")
147
+ p.add_argument("--thoughts-frac", type=float, default=0.25)
148
+ p.add_argument("--thoughts-dataset", default="volcanos/OpenThoughts2-1M-ShortThink")
149
+ p.add_argument("--thoughts-config", default=None)
150
+ p.add_argument("--thoughts-split", default="train")
151
+ p.add_argument("--seed", type=int, default=0)
152
+ args = p.parse_args()
153
+
154
+ torch.manual_seed(args.seed)
155
+ dev = args.device
156
+ tok = T.load_tokenizer(args.tok_dir)
157
+ pad = tok.pad_token_id or 0
158
+ render = make_renderer(tok)
159
+
160
+ ck = torch.load(args.init_ckpt, map_location=dev, weights_only=False)
161
+ cfg = QuazimotoConfig(**ck["family_config"])
162
+ model = QuazimotoLM(cfg); model.load_state_dict(ck["model"], strict=False)
163
+ model.to(dev)
164
+ print(f"SFT from {args.init_ckpt} (step {ck.get('step')}) | block {args.block} | UltraChat")
165
+
166
+ sources = [
167
+ (args.ultrachat_dataset, args.ultrachat_config, args.ultrachat_split, args.ultrachat_frac),
168
+ (args.feedback_dataset, args.feedback_config, args.feedback_split, args.feedback_frac),
169
+ (args.ultradata_dataset, args.ultradata_config, args.ultradata_split, args.ultradata_frac),
170
+ (args.thoughts_dataset, args.thoughts_config, args.thoughts_split, args.thoughts_frac),
171
+ ]
172
+ sources = [s for s in sources if s[3] > 0]
173
+ convos = blended_convos(sources, args.seed)
174
+
175
+ def sft_batch():
176
+ rows = []
177
+ while len(rows) < args.batch:
178
+ ids, mask = render(next(convos), args.block)
179
+ if sum(mask) > 0: # keep only convos with an assistant turn
180
+ rows.append((ids, mask))
181
+ L = max(len(r[0]) for r in rows)
182
+ x = torch.full((args.batch, L), pad, dtype=torch.long)
183
+ y = torch.full((args.batch, L), -1, dtype=torch.long) # -1 = ignore
184
+ for j, (ids, mask) in enumerate(rows):
185
+ t = torch.tensor(ids)
186
+ x[j, :len(ids)] = t
187
+ # predict token i+1 at position i; loss only where target is assistant content
188
+ for i in range(len(ids) - 1):
189
+ if mask[i + 1]:
190
+ y[j, i] = ids[i + 1]
191
+ return x.to(dev), y.to(dev)
192
+
193
+ opt = torch.optim.AdamW(model.parameters(), lr=args.lr, betas=(0.9, 0.95),
194
+ weight_decay=args.weight_decay)
195
+
196
+ start_step = 1
197
+ if args.resume and os.path.isfile(args.out):
198
+ r = torch.load(args.out, map_location=dev, weights_only=False)
199
+ model.load_state_dict(r["model"], strict=False)
200
+ if "optim" in r:
201
+ try: opt.load_state_dict(r["optim"])
202
+ except ValueError: pass
203
+ start_step = int(r.get("step", 0)) + 1
204
+ print(f"resumed SFT from {args.out} at step {r.get('step')} -> {start_step}")
205
+ if start_step > args.steps:
206
+ print(f" nothing to do: already at {start_step-1} >= --steps {args.steps}."); return
207
+
208
+ model.train()
209
+ t0 = time.time()
210
+ for step in range(start_step, args.steps + 1):
211
+ for g in opt.param_groups:
212
+ g["lr"] = T.lr_at(step, args.lr, args.warmup, args.steps, args.min_lr_frac)
213
+ x, y = sft_batch()
214
+ _, loss, aux = model(x, y)
215
+ total = loss + sum(aux.values())
216
+ opt.zero_grad(set_to_none=True)
217
+ total.backward()
218
+ gn = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
219
+ if torch.isfinite(gn):
220
+ opt.step()
221
+ else:
222
+ print(f"step {step}: non-finite grad, skip"); continue
223
+ if step % 10 == 0 or step == 1:
224
+ dt = time.time() - t0
225
+ ntok = int((y != -1).sum())
226
+ print(f"step {step:5d} | sft-loss {loss.item():.3f} | ppl {math.exp(min(loss.item(),20)):.1f} "
227
+ f"| asst-toks {ntok} | lr {opt.param_groups[0]['lr']:.2e} | {dt:.1f}s")
228
+ if args.ckpt_every and step % args.ckpt_every == 0:
229
+ T.save_ckpt(model, tok.vocab_size, step, args.out, opt)
230
+ T.save_ckpt(model, tok.vocab_size, args.steps, args.out, opt)
231
+ print("done ->", args.out)
232
+
233
+
234
+ if __name__ == "__main__":
235
+ main()
visualize.py ADDED
@@ -0,0 +1,287 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ visualize.py -- live visualizer for Mycel-LM: watch the fungal colony grow in 3D as
3
+ the model generates, token by token.
4
+
5
+ Captures per generated token (via instrument.Recorder, which MycelBlock writes to):
6
+ * tip POSITIONS (all pos_dim axes) for one layer's colony -- the 3D scatter
7
+ * per-tip growth TRAIL over the growth steps -- the hyphae extending this token
8
+ * per-layer colony density + trait-station sites + attention + token stream
9
+
10
+ Renders one self-contained HTML dashboard (data embedded -> opens via file://). The
11
+ colony is a Three.js network you can drag-rotate / scroll-zoom: hyphal tips linked to
12
+ their nearest neighbours as filaments (so it reads as a mycelial web, not loose dots),
13
+ tips coloured by local density, with trait STATIONS as orange wire-spheres and faint
14
+ grey growth trails. Three.js loads from a CDN (needs network the first time).
15
+
16
+ Usage:
17
+ python visualize.py --prompt "the mycelium spreads" --tokens 50
18
+ """
19
+ import argparse, json, os, sys, webbrowser
20
+ import torch
21
+
22
+ from model import QuazimotoLM, QuazimotoConfig
23
+ import instrument
24
+
25
+ PKG_DIR = os.path.dirname(os.path.abspath(__file__))
26
+
27
+
28
+ def find_ckpt(path):
29
+ if path and os.path.isfile(path):
30
+ return path
31
+ folder = (os.path.dirname(path) if path else os.path.join(PKG_DIR, "chkpt")) or "."
32
+ if not os.path.isdir(folder):
33
+ return None
34
+ pts = [os.path.join(folder, f) for f in os.listdir(folder) if f.endswith(".pt")]
35
+ return max(pts, key=os.path.getmtime) if pts else None
36
+
37
+
38
+ def load_tokenizer(tok_dir):
39
+ sys.path.insert(0, tok_dir)
40
+ from spike_tokenizer import SpikeTokenizer
41
+ return SpikeTokenizer(vocab_file=os.path.join(tok_dir, "tokenizer.json"))
42
+
43
+
44
+ @torch.no_grad()
45
+ def run_capture(model, cfg, tok, ids, n_tokens, temperature, top_k, device):
46
+ rec = instrument.Recorder(phase_layer=cfg.n_layer // 2, attn_layer=cfg.n_layer // 2)
47
+ instrument.set_rec(rec)
48
+ idx = torch.tensor([ids], device=device)
49
+ try:
50
+ for _ in range(n_tokens):
51
+ cond = idx[:, -cfg.block_size:]
52
+ rec.begin()
53
+ logits, _, _ = model(cond)
54
+ lg = torch.nan_to_num(logits[:, -1, :].float(), nan=0.0,
55
+ posinf=1e4, neginf=-1e4) / max(temperature, 1e-6)
56
+ if top_k:
57
+ v = torch.topk(lg, min(top_k, lg.size(-1))).values
58
+ lg = lg.masked_fill(lg < v[:, [-1]], float("-inf"))
59
+ probs = torch.softmax(lg, dim=-1)
60
+ nxt = (torch.argmax(lg, -1, keepdim=True) if temperature <= 1e-3
61
+ else torch.multinomial(probs, 1))
62
+ top = torch.topk(torch.softmax(logits[:, -1].float(), -1), 5)
63
+ rec.end(token=int(nxt),
64
+ char=tok.decode([int(nxt)], skip_special_tokens=False),
65
+ top=[[int(t), round(float(p), 3)] for t, p in zip(top.indices[0], top.values[0])])
66
+ idx = torch.cat([idx, nxt], dim=1)
67
+ finally:
68
+ instrument.set_rec(None)
69
+ return rec.frames, idx[0].tolist()
70
+
71
+
72
+ def main():
73
+ p = argparse.ArgumentParser(description="Mycel-LM 3D colony visualizer")
74
+ p.add_argument("--ckpt", default="")
75
+ p.add_argument("--tok_dir", default=PKG_DIR)
76
+ p.add_argument("--prompt", default="the mycelium spreads")
77
+ p.add_argument("--tokens", type=int, default=50)
78
+ p.add_argument("--temperature", type=float, default=0.0)
79
+ p.add_argument("--top_k", type=int, default=40)
80
+ p.add_argument("--device", default="cpu")
81
+ p.add_argument("--out", default=os.path.join(PKG_DIR, "viz.html"))
82
+ p.add_argument("--no_open", action="store_true")
83
+ args = p.parse_args()
84
+
85
+ path = find_ckpt(args.ckpt)
86
+ if path is None:
87
+ print("No checkpoint found; pass --ckpt or train one."); return
88
+ ckpt = torch.load(path, map_location=args.device, weights_only=False)
89
+ cfg = QuazimotoConfig(**ckpt["family_config"])
90
+ model = QuazimotoLM(cfg); model.load_state_dict(ckpt["model"], strict=False)
91
+ model.to(args.device).eval()
92
+ tok = load_tokenizer(args.tok_dir)
93
+ print(f"loaded {path} (step {ckpt.get('step')}) | capturing {args.tokens} tokens ...")
94
+
95
+ ids = tok.encode(args.prompt, add_special_tokens=False)
96
+ frames, _ = run_capture(model, cfg, tok, ids, args.tokens,
97
+ args.temperature, args.top_k or None, args.device)
98
+
99
+ pl = cfg.n_layer // 2
100
+ blk = model.layers[pl].quaz
101
+ stations = blk.stations.anchors.detach().cpu().tolist() if getattr(blk, "use_stations", False) else []
102
+
103
+ data = {
104
+ "meta": {"ckpt": os.path.basename(path), "step": ckpt.get("step"),
105
+ "n_layer": cfg.n_layer, "n_tips": cfg.n_tips, "pos_dim": cfg.mycel_pos_dim,
106
+ "bound": cfg.osc_bound, "phase_layer": pl, "prompt": args.prompt,
107
+ "stations": stations,
108
+ "tip_rings": bool(getattr(cfg, "use_tip_rings", False)),
109
+ "tip_ring_size": getattr(cfg, "tip_ring_size", 0)},
110
+ "prompt_chars": [tok.decode([t], skip_special_tokens=False) for t in ids],
111
+ "frames": frames,
112
+ }
113
+ html = HTML_TEMPLATE.replace("/*__DATA__*/", json.dumps(data))
114
+ with open(args.out, "w", encoding="utf-8") as f:
115
+ f.write(html)
116
+ print(f"wrote {args.out} ({os.path.getsize(args.out)//1024} KB)")
117
+ if not args.no_open:
118
+ webbrowser.open("file://" + os.path.abspath(args.out))
119
+
120
+
121
+ HTML_TEMPLATE = r"""<!doctype html>
122
+ <html><head><meta charset="utf-8"><title>Mycel-LM live</title>
123
+ <script src="https://cdnjs.cloudflare.com/ajax/libs/three.js/r128/three.min.js"></script>
124
+ <script src="https://cdn.jsdelivr.net/npm/three@0.128.0/examples/js/controls/OrbitControls.js"></script>
125
+ <style>
126
+ :root{--bg:#0d1117;--panel:#161b22;--ink:#e6edf3;--mut:#8b949e;--ac:#58a6ff;--hot:#f0883e;--spore:#3fb950}
127
+ *{box-sizing:border-box} body{margin:0;background:var(--bg);color:var(--ink);font:13px/1.4 ui-monospace,Menlo,Consolas,monospace}
128
+ header{padding:10px 16px;background:var(--panel);border-bottom:1px solid #30363d;display:flex;gap:16px;align-items:center;flex-wrap:wrap}
129
+ h1{font-size:15px;margin:0;color:var(--spore)} .mut{color:var(--mut)}
130
+ #wrap{display:grid;grid-template-columns:320px 1fr 300px;gap:12px;padding:12px}
131
+ .panel{background:var(--panel);border:1px solid #30363d;border-radius:8px;padding:10px}
132
+ .panel h2{font-size:12px;margin:0 0 8px;color:var(--mut);text-transform:uppercase;letter-spacing:.5px}
133
+ #stream{line-height:1.9;max-height:120px;overflow:auto}
134
+ .tk{padding:1px 2px;border-radius:3px;cursor:pointer;white-space:pre}
135
+ .tk.prompt{color:var(--mut)} .tk.cur{background:var(--spore);color:#000} .tk:hover{background:#30363d}
136
+ .ctrl{display:flex;gap:10px;align-items:center} input[type=range]{width:260px}
137
+ button{background:#21262d;color:var(--ink);border:1px solid #30363d;border-radius:6px;padding:4px 10px;cursor:pointer}
138
+ button:hover{border-color:var(--spore)} svg{display:block;width:100%;height:auto}
139
+ #colony3d{width:100%;height:440px;border-radius:6px;overflow:hidden}
140
+ .bar{height:10px;background:#21262d;border-radius:5px;overflow:hidden;margin:2px 0}.bar>div{height:100%}
141
+ table{border-collapse:collapse;width:100%}td{padding:1px 4px}.lbl{color:var(--mut);width:52px}
142
+ .grid{display:grid;gap:3px}.cell{height:15px;border-radius:2px}.small{font-size:11px;color:var(--mut)}
143
+ </style></head><body>
144
+ <header><h1>Mycel-LM &middot; colony 3D</h1><span class="mut" id="meta"></span>
145
+ <span style="flex:1"></span>
146
+ <div class="ctrl"><button id="play">▶ play</button><input type="range" id="seek" min="0" value="0"><span id="pos" class="mut"></span></div>
147
+ </header>
148
+ <div id="wrap">
149
+ <div style="display:flex;flex-direction:column;gap:12px">
150
+ <div class="panel"><h2>Token stream</h2><div id="stream"></div></div>
151
+ <div class="panel"><h2>Trait activity (this token)</h2><div id="traits"></div></div>
152
+ <div class="panel"><h2>Top predictions</h2><div id="top"></div></div>
153
+ </div>
154
+ <div style="display:flex;flex-direction:column;gap:12px">
155
+ <div class="panel"><h2>Colony &mdash; layer <span id="pl"></span> (drag rotate, scroll zoom; wire-spheres = trait stations)</h2>
156
+ <div id="colony3d"></div></div>
157
+ <div class="panel"><h2>Colony density per layer (bright = dense/clustered)</h2><div id="heat"></div></div>
158
+ </div>
159
+ <div style="display:flex;flex-direction:column;gap:12px">
160
+ <div class="panel"><h2>Attention &mdash; layer <span id="al"></span> (last token &rarr; context)</h2><div id="attn"></div></div>
161
+ <div class="panel"><h2>Notes</h2><div class="small" id="notes"></div></div>
162
+ </div>
163
+ </div>
164
+ <script>
165
+ const D=/*__DATA__*/; const M=D.meta,F=D.frames,NL=M.n_layer,NT=M.n_tips,PD=M.pos_dim,BND=M.bound;
166
+ let i=0,playing=false,timer=null; const $=id=>document.getElementById(id);
167
+ $("meta").textContent=`${M.ckpt} · step ${M.step} · ${NT} tips · ${PD}D colony · ${NL} layers · "${M.prompt}"`;
168
+ $("pl").textContent=M.phase_layer; $("al").textContent=M.phase_layer;
169
+ const seek=$("seek"); seek.max=F.length-1;
170
+ const hue2=v=>`hsl(${(1-Math.max(0,Math.min(1,v)))*140},70%,${30+45*Math.max(0,Math.min(1,v))}%)`;
171
+
172
+ // hsl(green->orange) -> [r,g,b] in 0..1 for Three.js vertex colours (density mode)
173
+ function colRGB(v){v=Math.max(0,Math.min(1,v));const h=(1-v)*140/360,s=0.7,l=0.32+0.42*v;
174
+ const a=s*Math.min(l,1-l),f=n=>{const k=(n+h*12)%12;return l-a*Math.max(-1,Math.min(k-3,9-k,1));};return[f(0),f(8),f(4)];}
175
+ // ring PHASE -> [r,g,b]: hue = phase around the wheel, saturation/brightness = coherence r.
176
+ // Synchronized tips share a hue (colony turns one colour); desynced tips are a rainbow.
177
+ function phaseRGB(psi,r){const h=(psi/(2*Math.PI))+0.5,s=0.25+0.65*Math.max(0,Math.min(1,r)),l=0.5;
178
+ const a=s*Math.min(l,1-l),f=n=>{const k=(n+h*12)%12;return l-a*Math.max(-1,Math.min(k-3,9-k,1));};return[f(0),f(8),f(4)];}
179
+
180
+ // ---- Three.js colony ----
181
+ const KNN=3; // links per tip -> the mycelial web
182
+ const MAXE=NT*KNN; // max filament segments
183
+ let scene,camera,renderer,controls,geom,points,webGeom,web,trailGeom,trail;
184
+ function init3d(){
185
+ const el=$("colony3d"),W=el.clientWidth||600,H=440;
186
+ scene=new THREE.Scene(); scene.background=new THREE.Color(0x0d1117);
187
+ camera=new THREE.PerspectiveCamera(55,W/H,0.1,1000); camera.position.set(BND*2.4,BND*1.7,BND*2.4);
188
+ renderer=new THREE.WebGLRenderer({antialias:true}); renderer.setSize(W,H); el.appendChild(renderer.domElement);
189
+ controls=new THREE.OrbitControls(camera,renderer.domElement); controls.enableDamping=true; controls.target.set(0,0,0);
190
+ scene.add(new THREE.LineSegments(new THREE.EdgesGeometry(new THREE.BoxGeometry(2*BND,2*BND,2*BND)),
191
+ new THREE.LineBasicMaterial({color:0x30363d})));
192
+ // filament web: each tip linked to its nearest neighbours -> looks like hyphae, not dots
193
+ webGeom=new THREE.BufferGeometry();
194
+ webGeom.setAttribute('position',new THREE.BufferAttribute(new Float32Array(MAXE*2*3),3));
195
+ webGeom.setAttribute('color',new THREE.BufferAttribute(new Float32Array(MAXE*2*3),3));
196
+ web=new THREE.LineSegments(webGeom,new THREE.LineBasicMaterial({vertexColors:true,transparent:true,opacity:0.55}));
197
+ scene.add(web);
198
+ // growth trails: each tip's path over the growth steps (the hypha extending this token)
199
+ const TS=(F[0]&&F[0].phases&&F[0].phases.traj)?F[0].phases.traj.length:0;
200
+ if(TS>1){trailGeom=new THREE.BufferGeometry();
201
+ trailGeom.setAttribute('position',new THREE.BufferAttribute(new Float32Array(NT*(TS-1)*2*3),3));
202
+ trail=new THREE.LineSegments(trailGeom,new THREE.LineBasicMaterial({color:0x8b949e,transparent:true,opacity:0.35}));
203
+ scene.add(trail);}
204
+ geom=new THREE.BufferGeometry();
205
+ geom.setAttribute('position',new THREE.BufferAttribute(new Float32Array(NT*3),3));
206
+ geom.setAttribute('color',new THREE.BufferAttribute(new Float32Array(NT*3),3));
207
+ points=new THREE.Points(geom,new THREE.PointsMaterial({size:BND*0.07,vertexColors:true}));
208
+ scene.add(points);
209
+ (M.stations||[]).forEach(s=>{const mesh=new THREE.Mesh(new THREE.SphereGeometry(BND*0.06,10,10),
210
+ new THREE.MeshBasicMaterial({color:0xf0883e,wireframe:true}));
211
+ mesh.position.set(s[0]||0,s[1]||0,s[2]||0); scene.add(mesh);});
212
+ window.addEventListener('resize',()=>{const w=el.clientWidth||600;camera.aspect=w/H;camera.updateProjectionMatrix();renderer.setSize(w,H);});
213
+ (function loop(){requestAnimationFrame(loop);controls.update();renderer.render(scene,camera);})();
214
+ }
215
+ function updateColony(f){
216
+ if(!geom||!f.phases)return;
217
+ const th=f.phases.theta,tp=f.phases.tip_psi,tr=f.phases.tip_r;
218
+ const pos=geom.attributes.position.array,col=geom.attributes.color.array;
219
+ const P=[]; for(let t=0;t<NT;t++)P.push([th[t*PD],th[t*PD+1],PD>2?th[t*PD+2]:0]);
220
+ let colr;
221
+ if(tp){ // colour by ring phase (synchronization view)
222
+ colr=P.map((_,t)=>phaseRGB(tp[t],tr?tr[t]:0.6));
223
+ }else{ // colour by local density (no rings)
224
+ const r2=(BND*0.25)**2, dens=P.map(a=>{let d=0;P.forEach(b=>{const dx=a[0]-b[0],dy=a[1]-b[1],dz=a[2]-b[2];if(dx*dx+dy*dy+dz*dz<r2)d++;});return d;});
225
+ const mx=Math.max(...dens,1); colr=dens.map(d=>colRGB(d/mx));
226
+ }
227
+ for(let t=0;t<NT;t++){pos[t*3]=P[t][0];pos[t*3+1]=P[t][1];pos[t*3+2]=P[t][2];
228
+ const c=colr[t];col[t*3]=c[0];col[t*3+1]=c[1];col[t*3+2]=c[2];}
229
+ geom.attributes.position.needsUpdate=true; geom.attributes.color.needsUpdate=true;
230
+ // rebuild the filament web: connect each tip to its KNN nearest neighbours
231
+ if(webGeom){const wp=webGeom.attributes.position.array,wc=webGeom.attributes.color.array;let e=0;
232
+ for(let a=0;a<NT&&e<MAXE;a++){
233
+ const nb=[]; for(let b=0;b<NT;b++){if(b===a)continue;
234
+ const dx=P[a][0]-P[b][0],dy=P[a][1]-P[b][1],dz=P[a][2]-P[b][2];nb.push([dx*dx+dy*dy+dz*dz,b]);}
235
+ nb.sort((x,y)=>x[0]-y[0]);
236
+ for(let n=0;n<Math.min(KNN,nb.length)&&e<MAXE;n++){const b=nb[n][1],ca=colr[a],cb=colr[b];
237
+ wp[e*6]=P[a][0];wp[e*6+1]=P[a][1];wp[e*6+2]=P[a][2];
238
+ wp[e*6+3]=P[b][0];wp[e*6+4]=P[b][1];wp[e*6+5]=P[b][2];
239
+ wc[e*6]=ca[0];wc[e*6+1]=ca[1];wc[e*6+2]=ca[2];wc[e*6+3]=cb[0];wc[e*6+4]=cb[1];wc[e*6+5]=cb[2];e++;}}
240
+ webGeom.setDrawRange(0,e*2);webGeom.attributes.position.needsUpdate=true;webGeom.attributes.color.needsUpdate=true;}
241
+ // growth trails: draw each tip's path across the captured growth steps
242
+ if(trail&&f.phases.traj){const tj=f.phases.traj,TS=tj.length,tp2=trailGeom.attributes.position.array;let g=0;
243
+ for(let t=0;t<NT;t++)for(let s=0;s<TS-1;s++){
244
+ const A=tj[s],Bp=tj[s+1];
245
+ tp2[g++]=A[t*PD];tp2[g++]=A[t*PD+1];tp2[g++]=PD>2?A[t*PD+2]:0;
246
+ tp2[g++]=Bp[t*PD];tp2[g++]=Bp[t*PD+1];tp2[g++]=PD>2?Bp[t*PD+2]:0;}
247
+ trailGeom.setDrawRange(0,NT*(TS-1)*2);trailGeom.attributes.position.needsUpdate=true;}
248
+ }
249
+
250
+ const stream=$("stream");
251
+ D.prompt_chars.forEach(c=>{const s=document.createElement("span");s.className="tk prompt";s.textContent=esc(c);stream.appendChild(s);});
252
+ F.forEach((f,k)=>{const s=document.createElement("span");s.className="tk gen";s.textContent=esc(f.char);s.onclick=()=>{i=k;render()};s.dataset.k=k;stream.appendChild(s);});
253
+ function esc(c){return c.replace(/\n/g,"⏎").replace(/ /g,"·");}
254
+ function heat(f){
255
+ let h=`<div class="grid" style="grid-template-columns:auto repeat(${NL},1fr)"><div class="small"></div>`;
256
+ for(let l=0;l<NL;l++)h+=`<div class="small" style="text-align:center">L${l}</div>`;
257
+ h+=`<div class="small">density</div>`; const mx=Math.max(...f.rings.map(r=>r.R[0]),1e-6);
258
+ for(let l=0;l<NL;l++){const s=f.rings[l]?f.rings[l].R[0]:0;h+=`<div class="cell" title="L${l} ${s.toFixed(3)}" style="background:${hue2(s/mx)}"></div>`;}
259
+ $("heat").innerHTML=h+"</div>";
260
+ }
261
+ function render(){
262
+ const f=F[i]; $("pos").textContent=`${i+1}/${F.length}`; seek.value=i;
263
+ document.querySelectorAll(".tk.gen").forEach(s=>s.classList.toggle("cur",+s.dataset.k===i));
264
+ updateColony(f); heat(f);
265
+ const qn=f.quaz_norm,mx=Math.max(...qn,1e-6);let L=qn.map((_,l)=>"grow L"+l),V=qn.map(v=>v/mx);
266
+ if(f.traits.hrm!=null){L.push("HRM");V.push(Math.min(1,f.traits.hrm/mx));}
267
+ if(f.traits.moe!=null){L.push("MoE");V.push(Math.min(1,f.traits.moe/mx));}
268
+ let th="<table>";V.forEach((v,k)=>{th+=`<tr><td class="lbl">${L[k]}</td><td style="width:100%"><div class="bar"><div style="width:${(v*100).toFixed(0)}%;background:${hue2(v)}"></div></div></td></tr>`;});
269
+ $("traits").innerHTML=th+"</table>";
270
+ let tp="<table>";f.top.forEach(([t,p])=>{tp+=`<tr><td class="lbl">${p.toFixed(2)}</td><td><div class="bar"><div style="width:${(p*100).toFixed(0)}%;background:var(--ac)"></div></div></td><td class="small">id ${t}</td></tr>`;});
271
+ $("top").innerHTML=tp+"</table>";
272
+ if(f.attn){const w=f.attn.w,m=Math.max(...w,1e-6);let h="<div style='display:flex;flex-wrap:wrap;gap:1px'>";
273
+ w.forEach((a,p)=>{h+=`<div title="pos ${p}: ${a.toFixed(3)}" style="width:8px;height:14px;background:${hue2(a/m)}"></div>`;});
274
+ $("attn").innerHTML=h+`</div><div class='small'>${w.length} context positions</div>`;}
275
+ else $("attn").innerHTML="<span class='small'>n/a</span>";
276
+ $("notes").innerHTML=`Tips: ${NT} · stations: ${(M.stations||[]).length}<br>Green = sparse growing edge, orange = dense core.<br>Filaments link each tip to its nearest neighbours (the mycelial web); faint grey trails are each tip's growth path this token.<br>Orange wire-spheres = trait stations. Drag to orbit, scroll to zoom.`;
277
+ }
278
+ seek.oninput=()=>{i=+seek.value;render()};
279
+ $("play").onclick=()=>{playing=!playing;$("play").textContent=playing?"⏸ pause":"▶ play";if(playing)timer=setInterval(()=>{i=(i+1)%F.length;render();},400);else clearInterval(timer);};
280
+ document.onkeydown=e=>{if(e.key==="ArrowRight"){i=Math.min(F.length-1,i+1);render();}if(e.key==="ArrowLeft"){i=Math.max(0,i-1);render();}};
281
+ if(window.THREE){init3d();} else {$("colony3d").innerHTML="<div class='small' style='padding:20px'>Three.js failed to load (needs network for the CDN).</div>";}
282
+ render();
283
+ </script></body></html>"""
284
+
285
+
286
+ if __name__ == "__main__":
287
+ main()