Add Mycel-LM 79M: architecture, tokenizer, training scripts, weights-only checkpoints
Browse files- README.md +317 -0
- build_fractal_table.py +106 -0
- chat_sft.py +129 -0
- chkpt/quazimoto.pt +3 -0
- chkpt/quazimoto_sft.pt +3 -0
- distill_uld.py +194 -0
- family.py +446 -0
- fractal.py +79 -0
- fractal_phase.pt +3 -0
- generate.py +326 -0
- healthcheck.py +258 -0
- instrument.py +84 -0
- model.py +773 -0
- mycel.py +178 -0
- opd_teacher.py +79 -0
- requirements.txt +11 -0
- special_tokens.py +85 -0
- spike_tokenizer.py +124 -0
- tokenizer.json +0 -0
- train.bat +24 -0
- train.py +297 -0
- train_opd.py +290 -0
- train_sft.bat +16 -0
- train_sft.py +235 -0
- 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 · 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 — 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 — layer <span id="al"></span> (last token → 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()
|