Upload LibraMind Mini weights, tokenizer, and inference code
Browse filesAdds model.safetensors (bf16), tokenizer.json, config.json, model.py, inference.py, chat_format.py, and the model card.
- README.md +326 -0
- chat_format.py +57 -0
- config.json +18 -0
- inference.py +335 -0
- model.py +263 -0
- model.safetensors +3 -0
- 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
|
|
|