alby13 commited on
Commit
bcda938
·
verified ·
1 Parent(s): 988837b

Upload LibraMind Mini weights, tokenizer, and inference code

Browse files

Adds model.safetensors (bf16), tokenizer.json, config.json, model.py, inference.py, chat_format.py, and the model card.

Files changed (7) hide show
  1. README.md +326 -0
  2. chat_format.py +57 -0
  3. config.json +18 -0
  4. inference.py +335 -0
  5. model.py +263 -0
  6. model.safetensors +3 -0
  7. tokenizer.json +0 -0
README.md CHANGED
@@ -1,3 +1,329 @@
1
  ---
2
  license: apache-2.0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3
  ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
  license: apache-2.0
3
+ language:
4
+ - en
5
+ library_name: pytorch
6
+ pipeline_tag: text-generation
7
+ tags:
8
+ - text-generation
9
+ - conversational
10
+ - chat
11
+ - function-calling
12
+ - roleplay
13
+ - hybrid
14
+ - gated-deltanet
15
+ - linear-attention
16
+ - muon-optimizer
17
+ - research
18
+ - open-weights
19
+ # datasets:
20
+ # - HuggingFaceFW/fineweb
21
+ # - HuggingFaceTB/smoltalk
22
  ---
23
+
24
+ # LibraMind (565M parameters)
25
+
26
+ **LibraMind** is a 565M-parameter chat model pretrained from scratch by [alby13](https://github.com/alby13/LibraMind-AI) on a single consumer GPU (NVIDIA RTX 4090). It is a hybrid architecture that interleaves **18 Gated DeltaNet** linear-recurrent layers with **6 Gated Attention** layers (a repeating `DDDA × 6` pattern), with Block Attention Residuals and logit soft-capping, trained with the Muon optimizer. It went through a full five-stage pipeline: pretraining, midtraining, supervised fine-tuning, preference tuning, and RL with verifiable rewards.
27
+
28
+
29
+ > **Research model: no safety training.**
30
+ > LibraMind has not undergone safety training or red-teaming. It can produce inaccurate, biased, offensive, or harmful text, and it will not reliably refuse requests. It also makes factual and arithmetic mistakes (see the examples below). It is released as a research artifact, provided "as is" without warranty of any kind, and the authors accept no liability for its outputs or for any use of the model. You are solely responsible for evaluating it and for any safeguards, filtering, and human oversight your use case needs. Do not deploy it in user-facing, high-stakes, or safety-critical settings without your own safety work.
31
+ > This notice is guidance. It does not add to or modify the terms of the Apache-2.0 license.
32
+
33
+ ## Model Summary
34
+
35
+ | | |
36
+ |---|---|
37
+ | **Developer** | alby13 |
38
+ | **Parameters** | 565M (565.2M) |
39
+ | **Model type** | Chat model |
40
+ | **Architecture** | Hybrid Gated DeltaNet + Gated Attention, decoder-only |
41
+ | **Layers / width** | 24 layers, hidden size 1,024 (`DDDA × 6`: 3 DeltaNet layers then 1 Attention layer, repeated 6 times) |
42
+ | **Vocabulary** | 32,768 (custom byte-level BPE) |
43
+ | **Context length** | 4,096 tokens (pretrained at 2,048, extended during midtraining) |
44
+ | **Weights precision** | bf16 (about 1.1 GB) |
45
+ | **Language** | English |
46
+ | **License** | Apache-2.0 |
47
+ | **Release** | 10/6/2026 |
48
+
49
+ ## Intended use
50
+
51
+ - Light assistant tasks: casual conversation, rewriting and summarizing short texts, formatted answers such as lists, sections and word limits, and simple function calling.
52
+ - Hobby and research: a fully documented small-model training run (architecture, data, every stage, evaluations).
53
+ - Game NPCs and interactive characters: dialogue driven by a character card in the system prompt, with optional schema-forced JSON for emotions and actions.
54
+
55
+ ## Intended Use
56
+
57
+ **Intended for:**
58
+ - Research on small hybrid recurrent/attention language models
59
+ - Studying low-resource, single-GPU training pipelines (pretraining through RL) and the Muon optimizer at small scale
60
+ - Experiments with on-device chat, role-play, and structured output (JSON / tool calls), with your own validation of the results (for example, validate JSON and compute numbers in code)
61
+
62
+ **Not intended for:**
63
+ - High-stakes or safety-critical decisions (medical, legal, financial, etc.)
64
+ - Unsupervised, user-facing deployment: the model has no safety training
65
+ - Use as a source of factual information without independent verification
66
+
67
+ ## How to Use
68
+
69
+ LibraMind uses a custom architecture, so it doesn't load in transformers, llama.cpp or other GGUF runtimes. It needs this repository's model.py, inference.py and chat_format.py, plus:
70
+
71
+ PyTorch with CUDA (trained with 2.14)
72
+ flash-linear-attention 0.5.2, with Triton
73
+ tokenizers
74
+ llguidance (optional, for grammar-forced JSON)
75
+ In this repository, the weights are runs/rl1/model_final.pt (float32, 2.26 GB) and the tokenizer is data/chat/sft/tokenizer.json.
76
+
77
+ ```python
78
+ import torch
79
+ from tokenizers import Tokenizer
80
+
81
+ from chat_format import ChatEncoder
82
+ from inference import GrammarFactory, generate
83
+ from model import LM, ModelConfig
84
+
85
+ ck = torch.load("model_final.pt", map_location="cpu", weights_only=False)
86
+ model = LM(ModelConfig(**{**ck["model_config"], "grad_ckpt": False}))
87
+ model.load_state_dict(ck["model"])
88
+ model = model.cuda().to(torch.bfloat16).eval()
89
+ enc = ChatEncoder(Tokenizer.from_file("tokenizer.json"))
90
+
91
+
92
+ def chat(messages, schema=None, max_new=300):
93
+ matcher = [GrammarFactory("tokenizer.json", enc.im_end, model.config.vocab_size).json(schema)] if schema else None
94
+ with torch.autocast("cuda", dtype=torch.bfloat16):
95
+ out = generate(model, [enc.encode_prompt(messages)], max_new, temperature=0.7, top_p=0.9, rep_penalty=1.1,
96
+ stop_ids=[enc.im_end], matchers=matcher, vocab=enc.tok.get_vocab_size())[0]
97
+ return enc.tok.decode(out)
98
+
99
+
100
+ print(chat([{"role": "system", "content": "You are Brom, a gruff dwarven blacksmith. Stay in character."},
101
+ {"role": "user", "content": "Can you fix my sword?"}]))
102
+ ```
103
+
104
+ generate uses cached decoding: Gated DeltaNet recurrent state plus an attention KV cache, replayed as a CUDA graph. It accepts many prompts at once, at about 1,050 tokens/s across 64 parallel conversations. The repository also has a browser chat window (chatui.cmd) and a console chat (chat.cmd).
105
+
106
+ Recommended settings:
107
+
108
+ - temperature 0.4, top-p 0.9, repetition penalty 1.1
109
+ - System prompt:
110
+ ````
111
+ "You are LibraMind, a friendly and thoughtful conversation partner made by alby13. Keep replies natural and to the point: a few sentences for casual chat, more detail only when it's needed. Ask a follow-up question when it helps the conversation.
112
+ ````
113
+ - temperature 0 for JSON and tool calls
114
+ - always pass stop_ids=[<|im_end|>]
115
+
116
+ **Requirements:**
117
+
118
+ **Memory:** about 1.1 GB for the weights in bf16, plus activations. The recurrent layers use constant-size state, and only the 6 attention layers grow a KV cache with context length.
119
+
120
+
121
+ ## Chat format
122
+
123
+ ChatML with a beginning-of-text token. Roles are system, user, assistant and tool:
124
+
125
+ ```
126
+ <|bos|><|im_start|>system
127
+ {system prompt}<|im_end|>
128
+ <|im_start|>user
129
+ {message}<|im_end|>
130
+ <|im_start|>assistant
131
+ {reply}<|im_end|>
132
+ ```
133
+
134
+ The system prompt is optional. Without one, the model behaves as the LibraMind assistant.
135
+
136
+
137
+ ## Characters and Personas
138
+
139
+ Put the character card in the system prompt: name, setting, personality, how they speak, and any facts they must get right. That means prices, quest details, names and places. Roleplay training used this pattern:
140
+
141
+ ```
142
+ You are {name}. Stay in character: speak as {name} would, in the first person, and never say you are an AI.
143
+
144
+ Character profile:
145
+ {profile}
146
+ ```
147
+
148
+ For game actions, force the reply into a JSON schema with GrammarFactory.json(schema). The output is then guaranteed to parse and match the schema. Example schema: examples/npc_schema.json, with fields line, emotion and action. Describe the fields in the system prompt too, so the content makes sense as well as the format.
149
+
150
+ ## Function calling
151
+
152
+ ```
153
+ List the tools in the system prompt in the format the model was trained on:
154
+
155
+ You can call the following functions:
156
+ <tools>
157
+ [{"type": "function", "function": {"name": "get_weather", "description": "...", "parameters": {...}}}]
158
+ </tools>
159
+
160
+ To call a function, reply with a JSON object inside <tool_call></tool_call> tags, for example:
161
+ <tool_call>
162
+ {"name": "function_name", "arguments": {"arg": "value"}}
163
+ </tool_call>
164
+ You may make several calls. Function results come back in <tool_response> tags.
165
+ ```
166
+
167
+ GrammarFactory.tool_calls(tools) can force calls to name a listed function, with arguments that match its schema.
168
+
169
+
170
+ ## Example Outputs
171
+
172
+ sampling settings used (temperature 0.4, top-p 0.9, repetition penalty 1.1).
173
+
174
+ **General chat**
175
+
176
+ > **User:** Hi! Who are you?
177
+ > **LibraMind:** I'm LibraMind, a small AI language model trained from scratch by alby13. How can I help?
178
+
179
+ **Role-play**
180
+
181
+ > **System Prompt:** You are Grom, a gruff dwarven blacksmith in the mountain town of Bringham. Stay in character and keep replies short.
182
+ > **User:** Where can I find some armor? I need some that also protects against magic.
183
+ > **LibraMind:** Well, I ain't no scholar, but I've heard tales from travelers. You might find some armor in the blacksmith's forge. Ask the old sage, Thorne. He knows all about such things.
184
+
185
+ > **System Prompt:** You are Jin, the stern innkeeper of the Gilded Redbird inn in the village of Ashfield. A room costs 15 silver a night and a bowl of stew costs 3 silver. (...)
186
+ > **User:** How much for a room tonight and some supper?
187
+ > **LibraMind:** A room for 15 silver is a good price. A bowl of stew for 3 silver will be 9 silver.
188
+
189
+
190
+
191
+ **Structured output** (schema-forced decoding; the system prompt asks for JSON with `line`, `emotion`, and `action`)
192
+
193
+ > **System:** You are Grom (...). Reply only with JSON: "line" is what you say, "emotion" is how you feel, and "action" is what you do (open_shop if the player wants to buy or get something repaired).
194
+ > **User:** I'd like to buy a new axe.
195
+
196
+ ```
197
+ Sample 1: {"line": "axe", "emotion": "happy", "action": "open_shop"}
198
+ Sample 2: {"line": "You've come to the right place, friend. I've got just the thing for you. What's your name, and what kind of axe are you looking for?", "emotion": "neutral", "action": "open_shop"}
199
+ ```
200
+
201
+ **Instruction following**
202
+
203
+ > **User:** Give me three tips for memorizing for studying as a numbered list.
204
+ > **LibraMind:**
205
+ >1. Create a dedicated study space: Having a specific area for studying can help you focus and stay motivated.
206
+ >2. Use mnemonic devices: These are memory techniques that can help you remember information by associating it with something you already know or have in mind.
207
+ >3. Break your study sessions into smaller intervals: This can help maintain your concentration and prevent burnout.
208
+
209
+
210
+ ## Architecture
211
+
212
+ The layout follows Qwen3.5-style hybrid models: three linear-attention layers for every full-attention layer. The plain residual stream is replaced by Block Attention Residuals.
213
+
214
+ | Component | Details |
215
+ |---|---|
216
+ | **Gated DeltaNet** (18 layers) | Constant-memory linear recurrent state using the delta rule, with a short 1D convolution (kernel 4) and output gating. About 10.5M parameters per layer. |
217
+ | **Gated Attention** (6 layers, every 4th) | Grouped-query attention: 16 query heads × 128 dim, 4 shared KV heads, QK-Norm, RoPE (base 10,000), and a sigmoid output gate. About 7.3M parameters per layer. Only these 6 layers keep a KV cache. |
218
+ | **SwiGLU MLP** (all 24 layers) | Hidden width 3,584, about 11.0M parameters per layer. |
219
+ | **Block Attention Residuals** | 8 blocks of 6 sub-layers; each sub-layer attends over summaries of earlier blocks with a learned query. About 50K parameters in total. |
220
+ | **Normalization / bias** | Pre-RMSNorm without learnable scale; no bias terms anywhere. |
221
+ | **Embeddings / head** | Untied: 33.6M input embedding + 33.6M output head. |
222
+ | **Logit soft-cap** | `logits = 15 · tanh(logits / 15)` |
223
+ | **Tokenizer** | Byte-level BPE, 32,768 vocab, GPT-4-style regex splitting, about 4.7 bytes per token. Trained on 2B characters of web text. |
224
+
225
+ Parameter split: about 498.1M in the transformer body (24 layers) and 67.1M in embeddings and output head.
226
+
227
+ ## Training
228
+
229
+ | Stage | Data
230
+ |---|---|
231
+ | 1. Pretraining | Enhanced FineWeb text
232
+ | 2. Midtraining | 50% unseen FineWeb, 50% conversations
233
+ | 3. Supervised fine-tuning (SFT)
234
+ | 4. Preference tuning
235
+ | 5. RL with verifiable rewards (GRPO)
236
+
237
+
238
+ **Pretraining details:**
239
+ - **Optimizers:** Muon for the ~497M 2D matrix weights (LR 0.02, cautious weight decay 0.07 on a cosine schedule); AdamW (LR 0.001) for embeddings, output head, gates, and Block Attention Residual queries
240
+ - **Schedule:** 40-step warmup, constant, then linear decay to 5% over the last 65% of training
241
+ - **Precision / compute:** bf16, `torch.compile`, gradient clipping at 1.0, about 17.1 GB VRAM on a single RTX 4090
242
+
243
+ ## Evaluation
244
+
245
+ ### Pretrained base model
246
+
247
+ Zero-shot evaluation with [lighteval](https://github.com/huggingface/lighteval) in cloze form (accuracy normalized by length; Winogrande and LAMBADA use plain accuracy). The SmolLM2-360M column is from its model card, which reports the same style of evaluation. SmolLM2-360M saw about 4 trillion training tokens, roughly 400× more than LibraMind.
248
+
249
+ | Benchmark | LibraMind base (10.2B tokens) | SmolLM2-360M (4T tokens) |
250
+ |---|---|---|
251
+ | HellaSwag | 55.3 | 54.5 |
252
+ | ARC (Easy / Challenge) | 65.6 / 35.4 (average 50.5) | average 53.0 |
253
+ | PIQA | 73.9 | 71.7 |
254
+ | OpenBookQA | 31.2 | 37.4 |
255
+ | CommonsenseQA | 40.6 | 38.0 |
256
+ | Social IQa | 44.6 | — |
257
+ | Winogrande | 56.6 | 52.5 |
258
+ | MMLU (cloze) | 30.3 | 35.8 |
259
+ | LAMBADA (accuracy / perplexity) | 46.3 / 13.3 | — |
260
+
261
+ ### Chat model, by training stage
262
+
263
+ - **IFEval:** 541 prompts, using the lm-eval implementation.
264
+ - **Tool calling:** 200 held-out function-calling conversations from the SFT sources. "Call decision" means it correctly chose whether to call a tool. Numbers in parentheses use grammar-forced (constrained) decoding.
265
+ - **Validation loss:** 1,800 held-out conversations.
266
+
267
+ | | mid | SFT | preference | RL (released) |
268
+ |---|---|---|---|---|
269
+ | IFEval prompt-level strict / loose | 40.7 / 42.3 | 45.8 / 49.7 | 50.5 / 53.8 | 51.4 / 54.5 |
270
+ | IFEval instruction-level strict / loose | 54.1 / 56.2 | 57.6 / 61.0 | 63.0 / 66.0 | 64.0 / 66.7 |
271
+ | Tools: correct call decision | 99.5 | 99.0 | 99.5 | 99.5 |
272
+ | Tools: valid JSON | 99.0 | 98.5 | 99.5 | 99.5 |
273
+ | Tools: right function (grammar-forced) | 87.8 (88.8) | 91.3 (92.9) | 92.4 (92.9) | 92.4 (92.4) |
274
+ | Tools: exact arguments | 71.4 | 80.1 | 81.1 | 80.6 |
275
+ | Validation loss, conversations | 1.035 | 0.950 | 0.966 | 0.968 |
276
+
277
+ On 40 held-out role-play prompts, the released model never described itself as an AI. Its average reply on general prompts is 156 words (175 after SFT).
278
+
279
+ **IFEval caveat:** these scores flatter the model's general instruction-following. Its training data includes IFEval-style constraint exercises, and the RL stage rewarded the same 19 instruction checkers that IFEval scores (on different prompts). For reference, SmolLM2-360M-Instruct reports 41.0 on IFEval (the average of its prompt- and instruction-level scores). LibraMind's equivalent strict average is 57.7.
280
+
281
+ ## Limitations and Bias
282
+
283
+ - Facts: it often states wrong or invented facts confidently. For example, it described the sitcom Last of the Summer Wine as a science-fiction series. Don't use it as a source of information.
284
+ - Math: it fails at arithmetic beyond single digits (see above). Part of the cause is the tokenizer, which splits numbers into irregular 1–3-digit chunks (6340 → 6|34|0, 12452 → 124|52). Do calculations in code, or give the model a calculator tool.
285
+ - Code: generated code is rarely correct.
286
+ - Using provided information: it usually uses facts from the system prompt, but can garble them, especially numbers. In game use, validate anything that matters.
287
+ - Roleplay: it stays in character but can drift from the character's details, hedge ("my expertise lies elsewhere, but…") or invent backstory. Long conversations lose coherence, and the context limit is 4,096 tokens.
288
+ - Language: English only.
289
+ - Safety: there was no dedicated safety training or red-teaming. Any refusal behavior comes only from the fine-tuning and preference data. Like any web-trained model, it can produce biased, offensive or harmful text. Filter outputs before showing them to players or the public.
290
+
291
+ ## Training Data and Licenses
292
+
293
+ - **FineWeb:** released under the permissive ODC-By 1.0 (Open Data Commons Attribution) license.
294
+ - **SmolTalk, smol-smoltalk, and SmolTalk2:** Apache 2.0
295
+ - **GSM8K and MMLU:** MIT.
296
+
297
+ ## License
298
+
299
+ The model weights and code are released under the **Apache License 2.0**; see the `LICENSE` file. Apache-2.0 includes a disclaimer of warranty and a limitation of liability (Sections 7 and 8). Training data carries its own terms; see above.
300
+
301
+ ## Acknowledgements
302
+
303
+ LibraMind builds on public research and open tools:
304
+
305
+ - Gated DeltaNet (Yang, Kautz and Hatamizadeh, 2024) and the flash-linear-attention library
306
+ - The Qwen3.5 hybrid layout and gated attention
307
+ - Block Attention Residuals from the Kimi team at Moonshot AI
308
+ - The Muon optimizer (Keller Jordan, 2024)
309
+ - Anchored Preference Optimization (D'Oosterlinck et al., 2024)
310
+ - GRPO (DeepSeekMath, 2024)
311
+ - IFEval (Zhou et al., 2023)
312
+ - lighteval
313
+ - The SmolLM2/SmolLM3 data and recipes from Hugging Face, Tulu 3 from Ai2, GSM8K, and MMLU
314
+
315
+ ## Citation
316
+
317
+ ```bibtex
318
+ @misc{alby13_libramind_2026,
319
+ author = {alby13},
320
+ title = {LibraMind: A 565M Parameter Hybrid Gated DeltaNet-Attention Chat Model},
321
+ year = {2026},
322
+ publisher = {Local AI Research},
323
+ url = {https://github.com/alby13/LibraMind-AI}
324
+ }
325
+ ```
326
+
327
+ ## Contact
328
+
329
+ GitHub issues discussions or Hugging Face Discussions tab.
chat_format.py ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Chat template shared by data preparation, fine-tuning and the chat program (ChatML style).
2
+
3
+ <|bos|><|im_start|>system\\n{system}<|im_end|>\\n<|im_start|>user\\n{text}<|im_end|>\\n<|im_start|>assistant\\n{reply}<|im_end|>\\n
4
+
5
+ Messages are encoded piece by piece, so training and inference tokenize identically.
6
+ Loss mask: 1 on assistant reply tokens and the <|im_end|> that closes them (so the model learns to stop).
7
+ """
8
+ import numpy as np
9
+ from tokenizers import Tokenizer
10
+
11
+ BOS, IM_START, IM_END = "<|bos|>", "<|im_start|>", "<|im_end|>"
12
+ ROLES = ("system", "user", "assistant", "tool")
13
+
14
+
15
+ def chat_tokenizer(base_tokenizer_path):
16
+ """The pretraining tokenizer plus the chat special tokens (appended as new ids)."""
17
+ tok = Tokenizer.from_file(str(base_tokenizer_path))
18
+ tok.add_special_tokens([IM_START, IM_END])
19
+ return tok
20
+
21
+
22
+ class ChatEncoder:
23
+ def __init__(self, tok):
24
+ self.tok = tok
25
+ self.bos, self.im_start, self.im_end = (tok.token_to_id(t) for t in (BOS, IM_START, IM_END))
26
+ self.newline = tok.encode("\n", add_special_tokens=False).ids
27
+ self.headers = {r: tok.encode(f"{r}\n", add_special_tokens=False).ids for r in ROLES}
28
+
29
+ def encode_conversation(self, messages):
30
+ """messages: list of {"role", "content"}. Returns (ids uint16, loss mask uint8)."""
31
+ ids, mask = [self.bos], [0]
32
+ for m in messages:
33
+ head = [self.im_start] + self.headers[m["role"]]
34
+ body = self.tok.encode(m["content"], add_special_tokens=False).ids + [self.im_end]
35
+ train = 1 if m["role"] == "assistant" else 0
36
+ ids += head + body + self.newline
37
+ mask += [0] * len(head) + [train] * len(body) + [0] * len(self.newline)
38
+ return np.array(ids, dtype=np.uint16), np.array(mask, dtype=np.uint8)
39
+
40
+ def encode_prompt(self, messages):
41
+ """Conversation so far plus the assistant header, ready for generation."""
42
+ ids, _ = self.encode_conversation(messages)
43
+ return np.concatenate([ids, [self.im_start] + self.headers["assistant"]]).astype(np.int64)
44
+
45
+ def encode_reply(self, text):
46
+ """An assistant reply as the model would produce it: its tokens followed by <|im_end|>."""
47
+ return np.array(self.tok.encode(text, add_special_tokens=False).ids + [self.im_end], dtype=np.int64)
48
+
49
+
50
+ def decode_conversation(tok, ids):
51
+ """Turn a tokenized conversation back into [{"role", "content"}] messages."""
52
+ text = tok.decode([int(i) for i in ids], skip_special_tokens=False).replace(BOS, "")
53
+ msgs = []
54
+ for part in text.split(IM_START)[1:]:
55
+ role, _, body = part.partition("\n")
56
+ msgs.append({"role": role, "content": body.split(IM_END)[0]})
57
+ return msgs
config.json ADDED
@@ -0,0 +1,18 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "vocab_size": 32832,
3
+ "n_layer": 24,
4
+ "d_model": 1024,
5
+ "arch": "hybrid",
6
+ "attn_every": 4,
7
+ "n_head": 16,
8
+ "n_kv_head": 4,
9
+ "head_dim": 128,
10
+ "gdn_heads": 16,
11
+ "gdn_head_dim": 128,
12
+ "ffn_dim": 3584,
13
+ "attnres_blocks": 8,
14
+ "rope_theta": 10000.0,
15
+ "max_seq_len": 4096,
16
+ "softcap": 15.0,
17
+ "grad_ckpt": false
18
+ }
inference.py ADDED
@@ -0,0 +1,335 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Fast generation: cached incremental decoding, batched sampling, optional grammar constraints.
2
+
3
+ Each sequence's state is small and fixed per token: attention layers keep their keys/values, DeltaNet
4
+ layers keep their recurrent state (flash-linear-attention cache), and Block AttnRes only mixes within
5
+ a token, so it needs no cache. A prompt is processed once (prefill) and each new token then costs one
6
+ cheap step instead of re-running the whole conversation.
7
+
8
+ outs = generate(model, [prompt_ids, ...], max_new=200, stop_ids=[enc.im_end])
9
+ matcher = GrammarFactory(tokenizer_path, im_end_id).json(schema) # output guaranteed to parse
10
+ """
11
+ import torch
12
+ import torch.nn.functional as F
13
+
14
+ from model import DeltaNet, GatedAttention, attnres_mix, inv_rms, rms
15
+
16
+
17
+ class _GDNStates:
18
+ """The minimal cache interface flash-linear-attention layers expect (len / [] / update)."""
19
+
20
+ def __init__(self, n):
21
+ self.layers = [None] * n
22
+
23
+ def __len__(self):
24
+ return len(self.layers)
25
+
26
+ def __getitem__(self, i):
27
+ return self.layers[i]
28
+
29
+ def update(self, layer_idx, recurrent_state=None, conv_state=None, **_):
30
+ cur = self.layers[layer_idx]
31
+ if cur is None: # first call (prefill): keep the tensors
32
+ self.layers[layer_idx] = dict(recurrent_state=recurrent_state, conv_state=conv_state)
33
+ return
34
+ # later calls: write into the existing buffers, so a replayed CUDA graph carries the state forward
35
+ pairs = [(cur["recurrent_state"], recurrent_state)] + list(zip(cur["conv_state"] or (), conv_state or ()))
36
+ for dst, src in pairs:
37
+ if src is not None and dst.data_ptr() != src.data_ptr():
38
+ dst.copy_(src)
39
+
40
+
41
+ class Cache:
42
+ """Decoding state for a batch: attention keys/values (preallocated to `capacity`), DeltaNet states,
43
+ and each row's current length (rows may have different lengths)."""
44
+
45
+ def __init__(self, model, batch, capacity):
46
+ dev, dt = model.embed.weight.device, torch.bfloat16
47
+ self.pos = torch.zeros(batch, dtype=torch.long, device=dev)
48
+ self.max_pos, self.capacity = 0, capacity
49
+ self.kv = {j: (torch.zeros(batch, s.nkv, capacity, s.hd, device=dev, dtype=dt),
50
+ torch.zeros(batch, s.nkv, capacity, s.hd, device=dev, dtype=dt))
51
+ for j, s in enumerate(model.sublayers) if isinstance(s, GatedAttention)}
52
+ self.gdn = _GDNStates(len(model.sublayers))
53
+
54
+ @staticmethod
55
+ def stack(caches, extra):
56
+ """Combine single-sequence caches of different lengths into one batch with room for `extra` tokens."""
57
+ out = Cache.__new__(Cache)
58
+ out.pos = torch.cat([c.pos for c in caches])
59
+ out.max_pos = max(c.max_pos for c in caches)
60
+ out.capacity = out.max_pos + extra
61
+ out.kv = {}
62
+ for j, (K0, _) in caches[0].kv.items():
63
+ K = K0.new_zeros(len(caches), K0.size(1), out.capacity, K0.size(3))
64
+ V = torch.zeros_like(K)
65
+ for i, c in enumerate(caches):
66
+ n = c.max_pos
67
+ K[i, :, :n], V[i, :, :n] = c.kv[j][0][0, :, :n], c.kv[j][1][0, :, :n]
68
+ out.kv[j] = (K, V)
69
+ out.gdn = _GDNStates(len(caches[0].gdn))
70
+ for i, layer in enumerate(caches[0].gdn.layers):
71
+ if layer is not None:
72
+ conv = layer["conv_state"]
73
+ out.gdn.layers[i] = dict(
74
+ recurrent_state=torch.cat([c.gdn.layers[i]["recurrent_state"] for c in caches]),
75
+ conv_state=None if conv is None else tuple(
76
+ torch.cat([c.gdn.layers[i]["conv_state"][k] for c in caches]) for k in range(len(conv))))
77
+ return out
78
+
79
+
80
+ def _attention(sub, x, K, V, pos):
81
+ """GatedAttention over cached keys/values; x holds t new tokens per row starting at position pos[row]."""
82
+ B, t, _ = x.shape
83
+ q, gate = sub.q_proj(x).split(sub.nh * sub.hd, -1)
84
+ q = q.view(B, t, sub.nh, sub.hd)
85
+ k, v = sub.kv_proj(x).view(B, t, 2, sub.nkv, sub.hd).unbind(2)
86
+ positions = pos[:, None] + torch.arange(t, device=x.device) # [B, t]
87
+ cos, sin = (b[0, :, 0][positions].unsqueeze(2).to(x.dtype) for b in (sub.cos, sub.sin))
88
+
89
+ def rope(z):
90
+ z1, z2 = z.chunk(2, -1)
91
+ return torch.cat([z1 * cos - z2 * sin, z1 * sin + z2 * cos], -1)
92
+
93
+ q, k = rope(rms(q)), rope(rms(k))
94
+ rows = torch.arange(B, device=x.device)[:, None].expand(B, t)
95
+ K[rows, :, positions] = k.to(K.dtype)
96
+ V[rows, :, positions] = v.to(V.dtype)
97
+ keys = K.repeat_interleave(sub.nh // sub.nkv, dim=1)
98
+ vals = V.repeat_interleave(sub.nh // sub.nkv, dim=1)
99
+ mask = torch.arange(K.size(2), device=x.device)[None, None, :] <= positions[:, :, None] # causal per row
100
+ y = F.scaled_dot_product_attention(q.transpose(1, 2), keys.to(q.dtype), vals.to(q.dtype), attn_mask=mask[:, None])
101
+ y = y.transpose(1, 2).reshape(B, t, sub.nh * sub.hd)
102
+ return sub.o_proj(y * torch.sigmoid(gate))
103
+
104
+
105
+ @torch.no_grad()
106
+ def step(model, idx, cache):
107
+ """Feed t new tokens per row (idx [B, t]); returns next-token logits [B, vocab] for the last position."""
108
+ t = idx.size(1)
109
+ if cache.max_pos + t > cache.capacity:
110
+ raise ValueError(f"cache full ({cache.capacity} tokens)")
111
+ x0 = rms(model.embed(idx))
112
+ blocks, blocks_inv, partial = [x0], [inv_rms(x0)], None
113
+ for j, sub in enumerate(model.sublayers):
114
+ srcs = blocks if partial is None else blocks + [partial]
115
+ inv = blocks_inv if partial is None else blocks_inv + [inv_rms(partial)]
116
+ h = rms(attnres_mix(srcs, inv, model.attnres_queries[j]))
117
+ if isinstance(sub, GatedAttention):
118
+ out = _attention(sub, h, *cache.kv[j], cache.pos)
119
+ elif isinstance(sub, DeltaNet):
120
+ sub.gdn.layer_idx = j
121
+ out = sub.gdn(h, past_key_values=cache.gdn, use_cache=True)[0]
122
+ else:
123
+ out = sub(h)
124
+ out = out.float()
125
+ partial = out if partial is None else partial + out
126
+ if (j + 1) % model.block_size == 0:
127
+ blocks.append(partial)
128
+ blocks_inv.append(inv_rms(partial))
129
+ partial = None
130
+ srcs = blocks if partial is None else blocks + [partial]
131
+ inv = blocks_inv if partial is None else blocks_inv + [inv_rms(partial)]
132
+ h = rms(attnres_mix(srcs, inv, model.attnres_queries[-1]))
133
+ cache.pos += t
134
+ cache.max_pos += t
135
+ return model._logits(h[:, -1])
136
+
137
+
138
+ def prefill(model, prompt, extra):
139
+ """Process one prompt (list/array of ids); returns (cache with room for `extra` tokens, last logits [1, V])."""
140
+ ids = torch.as_tensor(prompt, dtype=torch.long, device=model.embed.weight.device)[None]
141
+ cache = Cache(model, 1, ids.size(1) + extra)
142
+ return cache, step(model, ids, cache)
143
+
144
+
145
+ def sample(logits, temperature=0.7, top_p=0.9, top_k=0, recent=None, rep_penalty=1.0, bitmask=None, vocab=None):
146
+ """Pick one token per row. recent: [B, n] recently generated ids (-1 = none) for the repetition penalty."""
147
+ logits = logits.float()
148
+ if vocab is not None and vocab < logits.size(1):
149
+ logits[:, vocab:] = -float("inf") # padding rows of the embedding are not real tokens
150
+ if recent is not None and rep_penalty != 1.0:
151
+ seen = torch.zeros(logits.size(0), logits.size(1) + 1, device=logits.device, dtype=torch.bool)
152
+ seen.scatter_(1, torch.where(recent < 0, logits.size(1), recent), True)
153
+ seen = seen[:, :-1]
154
+ logits = torch.where(seen, torch.where(logits > 0, logits / rep_penalty, logits * rep_penalty), logits)
155
+ if bitmask is not None: # llguidance bitmask: bit i of word w allows token 32*w + i
156
+ bits = (bitmask.to(logits.device)[:, :, None] >> torch.arange(32, device=logits.device)) & 1
157
+ logits = logits.masked_fill(bits.reshape(bits.size(0), -1)[:, :logits.size(1)] == 0, -float("inf"))
158
+ if temperature <= 0:
159
+ return logits.argmax(-1)
160
+ logits = logits / temperature
161
+ if top_k:
162
+ kth = torch.topk(logits, top_k, dim=-1).values[:, -1:]
163
+ logits = logits.masked_fill(logits < kth, -float("inf"))
164
+ probs = torch.softmax(logits, -1)
165
+ if top_p < 1.0:
166
+ sp, si = probs.sort(-1, descending=True)
167
+ sp = sp.masked_fill(sp.cumsum(-1) - sp > top_p, 0.0)
168
+ probs = torch.zeros_like(probs).scatter_(-1, si, sp)
169
+ return torch.multinomial(probs, 1)[:, 0]
170
+
171
+
172
+ @torch.no_grad()
173
+ def generate(model, prompts, max_new=256, temperature=0.7, top_p=0.9, top_k=0, rep_penalty=1.0, stop_ids=(),
174
+ matchers=None, batch_size=32, vocab=None, on_token=None, cuda_graph=True):
175
+ """Sample continuations for many prompts, `batch_size` at a time. Returns a list of token-id lists.
176
+ matchers: optional per-prompt llguidance matchers (None = unconstrained) that force a grammar.
177
+ on_token(row_index, token_id): optional callback, e.g. for streaming a single prompt."""
178
+ from llguidance.torch import allocate_token_bitmask, fill_next_token_bitmask
179
+ results = [None] * len(prompts)
180
+ stop = torch.tensor(list(stop_ids) or [-1], device=model.embed.weight.device)
181
+ for start in range(0, len(prompts), batch_size):
182
+ idx = list(range(start, min(start + batch_size, len(prompts))))
183
+ pre = [prefill(model, prompts[i], max_new) for i in idx]
184
+ cache = Cache.stack([c for c, _ in pre], max_new) if len(pre) > 1 else pre[0][0]
185
+ logits = torch.cat([l for _, l in pre])
186
+ del pre
187
+ B = len(idx)
188
+ rows_m = [matchers[i] if matchers else None for i in idx]
189
+ bitmask = allocate_token_bitmask(B, logits.size(1)) if any(rows_m) else None
190
+ out = [[] for _ in range(B)]
191
+ done = torch.zeros(B, dtype=torch.bool, device=logits.device)
192
+ recent = torch.full((B, 64), -1, dtype=torch.long, device=logits.device)
193
+ graph = _DecodeGraph(model, cache, B) if cuda_graph else None
194
+ for n in range(max_new):
195
+ if bitmask is not None:
196
+ bitmask.fill_(-1) # all tokens allowed ...
197
+ for r, m in enumerate(rows_m):
198
+ if m is not None and not m.is_stopped():
199
+ fill_next_token_bitmask(m, bitmask, r) # ... except where a grammar forbids them
200
+ nxt = sample(logits, temperature, top_p, top_k, recent, rep_penalty, bitmask, vocab)
201
+ nxt = torch.where(done, stop[0].clamp(min=0), nxt)
202
+ toks = nxt.tolist()
203
+ for r in range(B):
204
+ if done[r]:
205
+ continue
206
+ if rows_m[r] is not None:
207
+ rows_m[r].consume_token(toks[r])
208
+ if toks[r] in stop_ids:
209
+ done[r] = True
210
+ continue
211
+ out[r].append(toks[r]) # keep it even if it completes the grammar (e.g. the final "}")
212
+ if on_token:
213
+ on_token(idx[r], toks[r])
214
+ if rows_m[r] is not None and rows_m[r].is_stopped():
215
+ done[r] = True
216
+ recent = torch.cat([recent[:, 1:], nxt[:, None]], 1)
217
+ if bool(done.all()) or n == max_new - 1:
218
+ break
219
+ logits = graph.step(nxt) if graph else step(model, nxt[:, None], cache)
220
+ for r, i in enumerate(idx):
221
+ results[i] = out[r]
222
+ return results
223
+
224
+
225
+ class _DecodeGraph:
226
+ """Records the one-token decode step as a CUDA graph and replays it: one launch instead of ~3,500
227
+ small kernel launches per token (launch overhead, not the GPU, dominates at this model size).
228
+ The first steps run normally (they also warm up the Triton kernels); if recording fails, it
229
+ quietly keeps running the normal way."""
230
+
231
+ WARMUP = 2
232
+
233
+ def __init__(self, model, cache, batch):
234
+ self.model, self.cache, self.n, self.graph = model, cache, 0, None
235
+ self.tok = torch.zeros(batch, 1, dtype=torch.long, device=model.embed.weight.device)
236
+ self.failed = False
237
+
238
+ def step(self, nxt):
239
+ self.n += 1
240
+ if self.failed or self.n <= self.WARMUP:
241
+ return step(self.model, nxt[:, None], self.cache)
242
+ if self.cache.max_pos + 1 > self.cache.capacity:
243
+ raise ValueError(f"cache full ({self.cache.capacity} tokens)")
244
+ self.tok.copy_(nxt[:, None])
245
+ if self.graph is None:
246
+ try:
247
+ self.graph = torch.cuda.CUDAGraph()
248
+ with torch.cuda.graph(self.graph), torch.autocast("cuda", dtype=torch.bfloat16, cache_enabled=False):
249
+ self.out = step(self.model, self.tok, self.cache) # recorded, not run
250
+ self.cache.max_pos -= 1 # the recording pass bumped it; the replay below is the real step
251
+ except Exception as e: # noqa: BLE001 - fall back to plain steps
252
+ print(f"(CUDA graph unavailable, using normal decoding: {type(e).__name__}: {str(e)[:120]})")
253
+ self.failed, self.graph = True, None
254
+ return step(self.model, nxt[:, None], self.cache)
255
+ self.graph.replay()
256
+ self.cache.max_pos += 1
257
+ return self.out
258
+
259
+
260
+ class GrammarFactory:
261
+ """Builds llguidance matchers for our tokenizer: guaranteed-valid JSON or tool calls."""
262
+
263
+ def __init__(self, tokenizer_path, stop_id, n_vocab):
264
+ """n_vocab: the model's (padded) output size, so masks line up with its logits."""
265
+ import llguidance.hf
266
+ from transformers import PreTrainedTokenizerFast
267
+ from tokenizers import Tokenizer
268
+ hf = PreTrainedTokenizerFast(tokenizer_object=Tokenizer.from_file(str(tokenizer_path)))
269
+ self.lltok = llguidance.hf.from_tokenizer(hf, n_vocab=n_vocab, eos_token=stop_id)
270
+
271
+ def _matcher(self, grammar):
272
+ from llguidance import LLMatcher
273
+ m = LLMatcher(self.lltok, grammar)
274
+ if m.is_error():
275
+ raise ValueError(m.get_error())
276
+ return m
277
+
278
+ COMPACT = {"whitespace_flexible": False, "item_separator": ", ", "key_separator": ": "}
279
+
280
+ def json(self, schema=None):
281
+ """Any valid JSON value matching `schema` (a dict; None = any JSON object). Compact whitespace, so a
282
+ reply can't wander off into blank space and run out of tokens before the JSON is closed."""
283
+ from llguidance import LLMatcher
284
+ return self._matcher(LLMatcher.grammar_from_json_schema({**(schema or {"type": "object"}), "x-guidance": self.COMPACT}))
285
+
286
+ def tool_calls(self, tools):
287
+ """One or more <tool_call> blocks whose JSON names a listed function with schema-valid arguments.
288
+ tools: OpenAI-style [{"type": "function", "function": {"name", "parameters"}}, ...]. Parameter lists in
289
+ the shorthand some datasets use ({"arg": {"type": "str, optional"}}) are converted to JSON Schema. If a
290
+ schema still can't be compiled, only the function name is enforced."""
291
+ from llguidance import LLMatcher
292
+ fns = [t.get("function", t) for t in tools]
293
+ try:
294
+ return self._tool_matcher(fns, strict_args=True)
295
+ except ValueError:
296
+ return self._tool_matcher(fns, strict_args=False)
297
+
298
+ def _tool_matcher(self, fns, strict_args):
299
+ import json as _json
300
+ from llguidance import LLMatcher
301
+ options = []
302
+ for fn in fns:
303
+ params = to_json_schema(fn.get("parameters")) if strict_args else {"type": "object"}
304
+ options.append({"type": "object", "properties": {"name": {"const": fn["name"]}, "arguments": params},
305
+ "required": ["name", "arguments"], "additionalProperties": False})
306
+ schema = _json.dumps({"anyOf": options, "x-guidance": self.COMPACT})
307
+ lark = (f'start: call ("\\n" call)*\n'
308
+ f'call: "<tool_call>\\n" body "\\n</tool_call>"\n'
309
+ f'body: %json {schema}\n')
310
+ return self._matcher(LLMatcher.grammar_from_lark(lark))
311
+
312
+
313
+ _SHORT_TYPES = {"str": "string", "string": "string", "int": "integer", "integer": "integer", "float": "number",
314
+ "number": "number", "bool": "boolean", "boolean": "boolean", "list": "array", "array": "array",
315
+ "dict": "object", "object": "object"}
316
+
317
+
318
+ def to_json_schema(params):
319
+ """Tool parameters as JSON Schema, accepting either real JSON Schema or the {"arg": {"type": "str, optional"}}
320
+ shorthand. Unknown types become 'any value'. Only declared arguments are allowed."""
321
+ if not params:
322
+ return {"type": "object", "properties": {}, "additionalProperties": False}
323
+ if params.get("type") == "object" or "properties" in params:
324
+ out = dict(params)
325
+ out.setdefault("additionalProperties", False)
326
+ return out
327
+ props, required = {}, []
328
+ for name, spec in params.items():
329
+ spec = spec if isinstance(spec, dict) else {}
330
+ raw = str(spec.get("type", "")).lower()
331
+ base = raw.split(",")[0].strip().split("[")[0]
332
+ props[name] = {"type": _SHORT_TYPES[base]} if base in _SHORT_TYPES else {}
333
+ if "optional" not in raw:
334
+ required.append(name)
335
+ return {"type": "object", "properties": props, "required": required, "additionalProperties": False}
model.py ADDED
@@ -0,0 +1,263 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Hybrid language model: Gated DeltaNet + gated attention (3:1), Block Attention Residuals.
2
+
3
+ Architecture follows Qwen3.5 (3 Gated DeltaNet layers per gated full-attention layer,
4
+ sigmoid output gate on attention, QK-norm) with Kimi's Block Attention Residuals replacing
5
+ the plain residual stream. arch="dense" swaps every DeltaNet layer for gated attention,
6
+ which gives a like-for-like Transformer baseline.
7
+ """
8
+ import math
9
+ from dataclasses import dataclass
10
+
11
+ import torch
12
+ import torch.nn as nn
13
+ import torch.nn.functional as F
14
+ from torch.utils.checkpoint import checkpoint
15
+
16
+
17
+ @dataclass
18
+ class ModelConfig:
19
+ vocab_size: int = 32768
20
+ n_layer: int = 24
21
+ d_model: int = 1024
22
+ arch: str = "hybrid" # "hybrid": every `attn_every`-th layer is attention, rest DeltaNet; "dense": all attention
23
+ attn_every: int = 4
24
+ n_head: int = 16 # attention query heads
25
+ n_kv_head: int = 4
26
+ head_dim: int = 128
27
+ gdn_heads: int = 16
28
+ gdn_head_dim: int = 128
29
+ ffn_dim: int = 3584
30
+ attnres_blocks: int = 8 # Block AttnRes: the 2*n_layer sublayers are split into this many blocks
31
+ rope_theta: float = 10000.0
32
+ max_seq_len: int = 2048
33
+ softcap: float = 15.0
34
+ grad_ckpt: bool = False # recompute sublayers in backward to save VRAM
35
+
36
+
37
+ PRESETS = {
38
+ # smoke tests only
39
+ "tiny": dict(n_layer=4, d_model=256, n_head=4, n_kv_head=2, head_dim=64, gdn_heads=4, gdn_head_dim=64,
40
+ ffn_dim=768, attnres_blocks=4),
41
+ # ~140M non-embedding: architecture bake-off size
42
+ "small": dict(n_layer=12, d_model=768, n_head=12, n_kv_head=3, head_dim=128, gdn_heads=12, gdn_head_dim=128,
43
+ ffn_dim=2688, attnres_blocks=8),
44
+ # ~500M non-embedding: Qwen3.5-0.8B layout with a 32k vocab
45
+ "base": dict(n_layer=24, d_model=1024, n_head=16, n_kv_head=4, head_dim=128, gdn_heads=16, gdn_head_dim=128,
46
+ ffn_dim=3584, attnres_blocks=8),
47
+ }
48
+
49
+
50
+ def rms(x):
51
+ return F.rms_norm(x, (x.size(-1),))
52
+
53
+
54
+ def inv_rms(x, eps=1e-6):
55
+ return torch.rsqrt(x.float().pow(2).mean(-1) + eps)
56
+
57
+
58
+ class GatedAttention(nn.Module):
59
+ """Causal GQA attention with QK-norm, RoPE and Qwen's elementwise sigmoid output gate."""
60
+
61
+ def __init__(self, c: ModelConfig):
62
+ super().__init__()
63
+ self.nh, self.nkv, self.hd = c.n_head, c.n_kv_head, c.head_dim
64
+ self.q_proj = nn.Linear(c.d_model, 2 * self.nh * self.hd, bias=False) # query and gate
65
+ self.kv_proj = nn.Linear(c.d_model, 2 * self.nkv * self.hd, bias=False)
66
+ self.o_proj = nn.Linear(self.nh * self.hd, c.d_model, bias=False)
67
+ inv_freq = 1.0 / (c.rope_theta ** (torch.arange(0, self.hd, 2).float() / self.hd))
68
+ freqs = torch.outer(torch.arange(c.max_seq_len).float(), inv_freq)
69
+ self.register_buffer("cos", freqs.cos()[None, :, None, :], persistent=False)
70
+ self.register_buffer("sin", freqs.sin()[None, :, None, :], persistent=False)
71
+
72
+ def rope(self, x):
73
+ T = x.size(1)
74
+ cos, sin = self.cos[:, :T].to(x.dtype), self.sin[:, :T].to(x.dtype)
75
+ x1, x2 = x.chunk(2, -1)
76
+ return torch.cat([x1 * cos - x2 * sin, x1 * sin + x2 * cos], -1)
77
+
78
+ def forward(self, x):
79
+ B, T, _ = x.shape
80
+ q, gate = self.q_proj(x).split(self.nh * self.hd, -1)
81
+ q = q.view(B, T, self.nh, self.hd)
82
+ k, v = self.kv_proj(x).view(B, T, 2, self.nkv, self.hd).unbind(2)
83
+ q, k = self.rope(rms(q)), self.rope(rms(k))
84
+ if self.nkv != self.nh:
85
+ k = k.repeat_interleave(self.nh // self.nkv, dim=2)
86
+ v = v.repeat_interleave(self.nh // self.nkv, dim=2)
87
+ y = F.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), is_causal=True)
88
+ y = y.transpose(1, 2).reshape(B, T, self.nh * self.hd)
89
+ return self.o_proj(y * torch.sigmoid(gate))
90
+
91
+
92
+ class DeltaNet(nn.Module):
93
+ """Gated DeltaNet (linear attention with a delta-rule memory) from flash-linear-attention."""
94
+
95
+ def __init__(self, c: ModelConfig):
96
+ super().__init__()
97
+ from fla.layers import GatedDeltaNet
98
+ self.gdn = GatedDeltaNet(hidden_size=c.d_model, head_dim=c.gdn_head_dim, num_heads=c.gdn_heads,
99
+ expand_v=1.0, mode="chunk", use_gate=True, use_short_conv=True, conv_size=4)
100
+
101
+ def forward(self, x):
102
+ return self.gdn(x)[0]
103
+
104
+
105
+ class SwiGLU(nn.Module):
106
+ def __init__(self, c: ModelConfig):
107
+ super().__init__()
108
+ self.w_gate = nn.Linear(c.d_model, c.ffn_dim, bias=False)
109
+ self.w_up = nn.Linear(c.d_model, c.ffn_dim, bias=False)
110
+ self.w_down = nn.Linear(c.ffn_dim, c.d_model, bias=False)
111
+
112
+ def forward(self, x):
113
+ return self.w_down(F.silu(self.w_gate(x)) * self.w_up(x))
114
+
115
+
116
+ def attnres_mix(sources, source_inv_rms, query):
117
+ """Block AttnRes: softmax over sources of <query, RMSNorm(source)>, then a weighted sum of sources."""
118
+ logits = torch.stack([(s @ query.to(s.dtype)).float() * r for s, r in zip(sources, source_inv_rms)])
119
+ weights = logits.softmax(0)
120
+ h = weights[0, ..., None] * sources[0]
121
+ for w, s in zip(weights[1:], sources[1:]):
122
+ h = h + w[..., None] * s
123
+ return h
124
+
125
+
126
+ class LM(nn.Module):
127
+ def __init__(self, c: ModelConfig):
128
+ super().__init__()
129
+ self.config = c
130
+ self.embed = nn.Embedding(c.vocab_size, c.d_model)
131
+ self.sublayers = nn.ModuleList()
132
+ for i in range(c.n_layer):
133
+ is_attn = c.arch == "dense" or (i % c.attn_every == c.attn_every - 1)
134
+ self.sublayers.append(GatedAttention(c) if is_attn else DeltaNet(c))
135
+ self.sublayers.append(SwiGLU(c))
136
+ n_sub = len(self.sublayers)
137
+ self.block_size = math.ceil(n_sub / c.attnres_blocks)
138
+ # one pseudo-query per sublayer plus one for the final read-out; zero init = uniform average at start
139
+ self.attnres_queries = nn.Parameter(torch.zeros(n_sub + 1, c.d_model))
140
+ self.lm_head = nn.Linear(c.d_model, c.vocab_size, bias=False)
141
+ self.init_weights()
142
+
143
+ @torch.no_grad()
144
+ def init_weights(self):
145
+ c = self.config
146
+ nn.init.normal_(self.embed.weight, std=1.0)
147
+ nn.init.normal_(self.lm_head.weight, std=0.001)
148
+ s = 3 ** 0.5 * c.d_model ** -0.5
149
+ for m in self.sublayers:
150
+ if isinstance(m, GatedAttention):
151
+ nn.init.uniform_(m.q_proj.weight, -s, s)
152
+ nn.init.uniform_(m.kv_proj.weight, -s, s)
153
+ nn.init.zeros_(m.o_proj.weight)
154
+ elif isinstance(m, SwiGLU):
155
+ nn.init.uniform_(m.w_gate.weight, -s, s)
156
+ nn.init.uniform_(m.w_up.weight, -s, s)
157
+ nn.init.zeros_(m.w_down.weight)
158
+ else: # DeltaNet keeps FLA's init for its gates/decay; match the rest to the attention layers
159
+ g = m.gdn
160
+ for lin in (g.q_proj, g.k_proj, g.v_proj, g.g_proj):
161
+ nn.init.uniform_(lin.weight, -s, s)
162
+ nn.init.zeros_(g.o_proj.weight)
163
+
164
+ def num_params(self, non_embedding=True):
165
+ n = sum(p.numel() for p in self.parameters())
166
+ return n - self.embed.weight.numel() - self.lm_head.weight.numel() if non_embedding else n
167
+
168
+ def _run(self, sub, x):
169
+ if self.config.grad_ckpt and self.training:
170
+ return checkpoint(sub, x, use_reentrant=False)
171
+ return sub(x)
172
+
173
+ def hidden(self, idx):
174
+ """Final normalized hidden states [B, T, d_model]."""
175
+ x0 = rms(self.embed(idx))
176
+ blocks, blocks_inv = [x0], [inv_rms(x0)] # completed block sums (block 0 = token embedding)
177
+ partial = None # running sum of sublayer outputs in the current block
178
+ for j, sub in enumerate(self.sublayers):
179
+ srcs = blocks if partial is None else blocks + [partial]
180
+ inv = blocks_inv if partial is None else blocks_inv + [inv_rms(partial)]
181
+ h = attnres_mix(srcs, inv, self.attnres_queries[j])
182
+ out = self._run(sub, rms(h)).float()
183
+ partial = out if partial is None else partial + out
184
+ if (j + 1) % self.block_size == 0:
185
+ blocks.append(partial)
186
+ blocks_inv.append(inv_rms(partial))
187
+ partial = None
188
+ srcs = blocks if partial is None else blocks + [partial]
189
+ inv = blocks_inv if partial is None else blocks_inv + [inv_rms(partial)]
190
+ return rms(attnres_mix(srcs, inv, self.attnres_queries[-1]))
191
+
192
+ def forward(self, idx, targets=None):
193
+ h = self.hidden(idx)
194
+ if targets is None:
195
+ return self._logits(h)
196
+ # Loss in chunks, recomputing each chunk's logits in backward, so the full
197
+ # (tokens x vocab) fp32 logit tensor never has to sit in VRAM.
198
+ # Targets of -1 (prompts, padding during fine-tuning) are ignored.
199
+ h, targets = h.flatten(0, 1), targets.flatten()
200
+ loss = 0.0
201
+ for i in range(0, h.size(0), self.LOSS_CHUNK):
202
+ loss = loss + checkpoint(self._chunk_loss, h[i:i + self.LOSS_CHUNK], targets[i:i + self.LOSS_CHUNK],
203
+ use_reentrant=False)
204
+ return loss / (targets >= 0).sum().clamp(min=1)
205
+
206
+ LOSS_CHUNK = 4096
207
+
208
+ def _logits(self, h):
209
+ logits = self.lm_head(h).float()
210
+ if self.config.softcap:
211
+ logits = self.config.softcap * torch.tanh(logits / self.config.softcap)
212
+ return logits
213
+
214
+ def _chunk_loss(self, h, targets):
215
+ return F.cross_entropy(self._logits(h), targets, reduction="sum", ignore_index=-1)
216
+
217
+ def token_logprobs(self, idx, targets):
218
+ """Log-probability of each target token, [B, T] (0 where the target is -1). Chunked like the loss."""
219
+ h, t = self.hidden(idx).flatten(0, 1), targets.flatten()
220
+ parts = [checkpoint(self._chunk_logprobs, h[i:i + self.LOSS_CHUNK], t[i:i + self.LOSS_CHUNK], use_reentrant=False)
221
+ for i in range(0, h.size(0), self.LOSS_CHUNK)]
222
+ return torch.cat(parts).view(targets.shape)
223
+
224
+ def _chunk_logprobs(self, h, targets):
225
+ lp = torch.log_softmax(self._logits(h), -1).gather(1, targets.clamp(min=0)[:, None])[:, 0]
226
+ return lp * (targets >= 0)
227
+
228
+ def param_groups(self):
229
+ """Split parameters: hidden matrices go to Muon, everything else to AdamW."""
230
+ muon, other = [], []
231
+ for name, p in self.sublayers.named_parameters():
232
+ (muon if p.ndim == 2 and min(p.shape) >= 64 else other).append(p)
233
+ other.append(self.attnres_queries)
234
+ return dict(muon=muon, embed=[self.embed.weight], head=[self.lm_head.weight], other=other)
235
+
236
+
237
+ def resize_vocab(state_dict, new_rows):
238
+ """Grow the embedding and output matrices for added tokens; new rows start at the mean of the old ones."""
239
+ for key in ("embed.weight", "lm_head.weight"):
240
+ w = state_dict[key]
241
+ if w.size(0) < new_rows:
242
+ extra = w.float().mean(0, keepdim=True).expand(new_rows - w.size(0), -1).to(w.dtype)
243
+ state_dict[key] = torch.cat([w, extra])
244
+ return state_dict
245
+
246
+
247
+ @torch.no_grad()
248
+ def generate(model, idx, max_new_tokens, temperature=0.8, top_k=50):
249
+ """Simple sampling that re-runs the full context each step (no cache; fine for short samples)."""
250
+ for _ in range(max_new_tokens):
251
+ ctx = idx[:, -model.config.max_seq_len:]
252
+ with torch.autocast("cuda", dtype=torch.bfloat16):
253
+ logits = model(ctx)[:, -1, :].float()
254
+ if temperature <= 0:
255
+ nxt = logits.argmax(-1, keepdim=True)
256
+ else:
257
+ logits = logits / temperature
258
+ if top_k:
259
+ v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
260
+ logits[logits < v[:, [-1]]] = -float("inf")
261
+ nxt = torch.multinomial(F.softmax(logits, -1), 1)
262
+ idx = torch.cat([idx, nxt], 1)
263
+ return idx
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:2e432bee4331517c77555a4b3e3b03e1148f26f2b6877f299a9e6355d7241290
3
+ size 1130734392
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff