Upload Haiku-base (pretrain EMA @ step 8400)
Browse files- README.md +275 -0
- __pycache__/inference.cpython-312.pyc +0 -0
- chat_template.jinja +86 -0
- config.json +94 -0
- inference.py +501 -0
- merges.txt +0 -0
- model.safetensors +3 -0
- requirements.txt +9 -0
- special_token_ids.json +10 -0
- special_tokens_map.json +68 -0
- tiny_gdn/__init__.py +9 -0
- tiny_gdn/config.py +193 -0
- tiny_gdn/haiku_layers.py +199 -0
- tiny_gdn/model.py +713 -0
- tiny_gdn/nn_common.py +33 -0
- tokenizer.json +0 -0
- tokenizer_config.json +74 -0
- validation.json +21 -0
- vocab.json +0 -0
- windows_fla_patches/fla/__init__.py +10 -0
- windows_fla_patches/fla/layers/__init__.py +63 -0
- windows_fla_patches/fla/ops/__init__.py +76 -0
- windows_fla_patches/fla/ops/simple_gla/__init__.py +23 -0
README.md
ADDED
|
@@ -0,0 +1,275 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: apache-2.0
|
| 3 |
+
language:
|
| 4 |
+
- en
|
| 5 |
+
tags:
|
| 6 |
+
- text-generation
|
| 7 |
+
- causal-lm
|
| 8 |
+
- pytorch
|
| 9 |
+
- pretrain
|
| 10 |
+
- hybrid
|
| 11 |
+
- kimi-delta-attention
|
| 12 |
+
- gated-mla
|
| 13 |
+
- haiku
|
| 14 |
+
pipeline_tag: text-generation
|
| 15 |
+
library_name: tiny_gdn
|
| 16 |
+
datasets:
|
| 17 |
+
- HuggingFaceFW/fineweb-edu
|
| 18 |
+
model-index:
|
| 19 |
+
- name: Haiku-base
|
| 20 |
+
results: []
|
| 21 |
+
---
|
| 22 |
+
|
| 23 |
+
<div align="center">
|
| 24 |
+
|
| 25 |
+
# Haiku-base
|
| 26 |
+
|
| 27 |
+
### Pretrained base model for the Haiku family (~655M)
|
| 28 |
+
|
| 29 |
+
[](.)
|
| 30 |
+
[-orange.svg)](.)
|
| 31 |
+
[](LICENSE)
|
| 32 |
+
[](.)
|
| 33 |
+
[](https://huggingface.co/spaces/kerzgrr/haiku-demo)
|
| 34 |
+
|
| 35 |
+
*A larger TinyGDN hybrid: Kimi Delta Attention memory plus gated multi-head latent attention*
|
| 36 |
+
|
| 37 |
+
</div>
|
| 38 |
+
|
| 39 |
+
---
|
| 40 |
+
|
| 41 |
+
## What this is
|
| 42 |
+
|
| 43 |
+
**Haiku-base** is the **pretrained (base) checkpoint** for **Haiku**, the ~655M successor to the Tercet family.
|
| 44 |
+
|
| 45 |
+
- Scales [`kerzgrr/Tercet-base`](https://huggingface.co/kerzgrr/Tercet-base) from ~502M to ~655M parameters
|
| 46 |
+
- Hybrid **Kimi Delta Attention (KDA)** recurrent layers + **gated MLA** (NoPE) full-attention layers
|
| 47 |
+
- Own **65,536** BPE tokenizer (not the Tercet 49k vocab)
|
| 48 |
+
- This repo is **pretrain-only** raw text continuation
|
| 49 |
+
- Chat / instruction SFT is **not released**
|
| 50 |
+
|
| 51 |
+
This base model is for continuation and research. It will not follow instructions reliably.
|
| 52 |
+
|
| 53 |
+
---
|
| 54 |
+
|
| 55 |
+
## Model Architecture
|
| 56 |
+
|
| 57 |
+
**Pipeline:** `Text Prompt` → `BPE-65K Tokenizer` → `Haiku Hybrid Decoder (36L)` → `Next-token Prediction`
|
| 58 |
+
|
| 59 |
+
### Hybrid block schedule (×36)
|
| 60 |
+
|
| 61 |
+
Every 4th layer is gated MLA; the rest are Kimi Delta Attention:
|
| 62 |
+
|
| 63 |
+
`KDA, KDA, KDA, MLA, …` (3:1 recurrent-to-attention)
|
| 64 |
+
|
| 65 |
+
| Component | Details |
|
| 66 |
+
|-----------|---------|
|
| 67 |
+
| **Kimi Delta Attention** | Linear-time recurrent memory (`flash-linear-attention`) |
|
| 68 |
+
| **Gated MLA** | DeepSeek-style latent KV, content-only QK (NoPE), full-rank output gate |
|
| 69 |
+
| **MLP** | SiTU-GLU |
|
| 70 |
+
| **Residuals** | Block attention residual |
|
| 71 |
+
| **Norm** | Zero-centered RMSNorm |
|
| 72 |
+
| **Embeddings** | Tied input / output |
|
| 73 |
+
|
| 74 |
+
### Technical specifications
|
| 75 |
+
|
| 76 |
+
| | |
|
| 77 |
+
|--|--|
|
| 78 |
+
| **Architecture** | Haiku hybrid (KDA + gated MLA) |
|
| 79 |
+
| **Parameters** | 655,270,488 deployable |
|
| 80 |
+
| **Hidden size** | 1,024 |
|
| 81 |
+
| **Intermediate (MLP)** | 3,840 |
|
| 82 |
+
| **Layers** | 36 (27 KDA + 9 gated MLA) |
|
| 83 |
+
| **Attention** | 8 heads, Q LoRA rank 512, KV LoRA rank 256 |
|
| 84 |
+
| **Linear (KDA)** | 8 heads × 128 dim |
|
| 85 |
+
| **Context (trained)** | 2,048 |
|
| 86 |
+
| **Max position embeddings** | 32,768 |
|
| 87 |
+
| **Vocabulary** | 65,536 (BPE) |
|
| 88 |
+
| **RoPE θ** | 1,000,000 (partial factor 0.5; used by KDA) |
|
| 89 |
+
| **Precision (Hub weights)** | bfloat16 EMA |
|
| 90 |
+
| **Weight file** | `model.safetensors` (~1.22 GiB) |
|
| 91 |
+
|
| 92 |
+
---
|
| 93 |
+
|
| 94 |
+
## Training (pretrain)
|
| 95 |
+
|
| 96 |
+
| | |
|
| 97 |
+
|--|--|
|
| 98 |
+
| **Dataset** | [FineWeb-Edu](https://huggingface.co/datasets/HuggingFaceFW/fineweb-edu) (10.13B packed train tokens) |
|
| 99 |
+
| **Tokens seen** | 4,404,019,200 |
|
| 100 |
+
| **Sequence length** | 2,048 |
|
| 101 |
+
| **Objective** | Next-token prediction (+ MTP during training; not used at decode) |
|
| 102 |
+
| **Optimizer** | Hybrid Muon + AdamW — β₁=0.9, β₂=0.95 |
|
| 103 |
+
| **Peak LR** | 2 × 10⁻⁴ |
|
| 104 |
+
| **Warmup** | 1% of steps |
|
| 105 |
+
| **Grad clip** | 1.0 |
|
| 106 |
+
| **EMA** | Karras power EMA (γ=1.0, p=0.75, max decay 0.9999) — **this Hub file is the EMA weights** |
|
| 107 |
+
| **Checkpoint** | optimizer step 8,400 |
|
| 108 |
+
| **Val loss (EMA)** | 3.6904 (ppl 40.06) |
|
| 109 |
+
|
| 110 |
+
---
|
| 111 |
+
|
| 112 |
+
## Install
|
| 113 |
+
|
| 114 |
+
### 1) System requirements
|
| 115 |
+
|
| 116 |
+
- Python **3.10+**
|
| 117 |
+
- **CUDA GPU strongly recommended**
|
| 118 |
+
- PyTorch with CUDA matching your driver
|
| 119 |
+
|
| 120 |
+
### 2) Create an environment
|
| 121 |
+
|
| 122 |
+
```bash
|
| 123 |
+
python -m venv .venv
|
| 124 |
+
# Windows
|
| 125 |
+
.venv\Scripts\activate
|
| 126 |
+
# Linux / macOS
|
| 127 |
+
source .venv/bin/activate
|
| 128 |
+
```
|
| 129 |
+
|
| 130 |
+
### 3) Install PyTorch
|
| 131 |
+
|
| 132 |
+
Pick the build for your platform from https://pytorch.org. Example:
|
| 133 |
+
|
| 134 |
+
```bash
|
| 135 |
+
pip install torch --index-url https://download.pytorch.org/whl/cu124
|
| 136 |
+
```
|
| 137 |
+
|
| 138 |
+
CPU-only:
|
| 139 |
+
|
| 140 |
+
```bash
|
| 141 |
+
pip install torch
|
| 142 |
+
```
|
| 143 |
+
|
| 144 |
+
### 4) Install Python deps
|
| 145 |
+
|
| 146 |
+
```bash
|
| 147 |
+
pip install safetensors tokenizers huggingface_hub
|
| 148 |
+
```
|
| 149 |
+
|
| 150 |
+
**Flash Linear Attention is installed automatically by `inference.py`** on first run (pinned commit + Windows import patches when needed). Git must be on `PATH`.
|
| 151 |
+
|
| 152 |
+
### 5) Download the inference script
|
| 153 |
+
|
| 154 |
+
```bash
|
| 155 |
+
curl -L -o inference.py https://huggingface.co/kerzgrr/Haiku-base/resolve/main/inference.py
|
| 156 |
+
|
| 157 |
+
# or Hugging Face CLI
|
| 158 |
+
hf download kerzgrr/Haiku-base inference.py --local-dir .
|
| 159 |
+
```
|
| 160 |
+
|
| 161 |
+
The script auto-downloads `model.safetensors`, `config.json`, `tokenizer.json`, and the `tiny_gdn/` package from this repo.
|
| 162 |
+
|
| 163 |
+
---
|
| 164 |
+
|
| 165 |
+
## Quick start
|
| 166 |
+
|
| 167 |
+
**Single prompt (streams tokens):**
|
| 168 |
+
|
| 169 |
+
```bash
|
| 170 |
+
python inference.py --prompt "The history of computing begins"
|
| 171 |
+
```
|
| 172 |
+
|
| 173 |
+
**Interactive REPL:**
|
| 174 |
+
|
| 175 |
+
```bash
|
| 176 |
+
python inference.py
|
| 177 |
+
```
|
| 178 |
+
|
| 179 |
+
**Common options:**
|
| 180 |
+
|
| 181 |
+
| Flag | Default | Description |
|
| 182 |
+
|------|---------|-------------|
|
| 183 |
+
| `--prompt` | *(none)* | One-shot continuation; omit for REPL |
|
| 184 |
+
| `--temperature` | `0.8` | Sampling temperature |
|
| 185 |
+
| `--top-p` | `0.95` | Nucleus sampling |
|
| 186 |
+
| `--top-k` | `50` | Top-k (0 disables) |
|
| 187 |
+
| `--max-new-tokens` | `256` | Generation length |
|
| 188 |
+
| `--repetition-penalty` | `1.08` | Repetition penalty |
|
| 189 |
+
| `--context-length` | `2048` | Tokens kept in the window |
|
| 190 |
+
| `--seed` | `42` | RNG seed |
|
| 191 |
+
| `--device` | `cuda` if available | `cuda` or `cpu` |
|
| 192 |
+
| `--no-stream` | off | Print the full completion at once |
|
| 193 |
+
| `--no-bos` | off | Do not prepend `<\|begin_of_text\|>` |
|
| 194 |
+
| `--local-dir` | *(none)* | Use a local snapshot directory |
|
| 195 |
+
|
| 196 |
+
---
|
| 197 |
+
|
| 198 |
+
## Files
|
| 199 |
+
|
| 200 |
+
```
|
| 201 |
+
kerzgrr/Haiku-base/
|
| 202 |
+
README.md
|
| 203 |
+
inference.py
|
| 204 |
+
requirements.txt
|
| 205 |
+
model.safetensors
|
| 206 |
+
config.json
|
| 207 |
+
tokenizer.json
|
| 208 |
+
tokenizer_config.json
|
| 209 |
+
special_tokens_map.json
|
| 210 |
+
special_token_ids.json
|
| 211 |
+
merges.txt
|
| 212 |
+
vocab.json
|
| 213 |
+
chat_template.jinja
|
| 214 |
+
tiny_gdn/
|
| 215 |
+
__init__.py
|
| 216 |
+
config.py
|
| 217 |
+
model.py
|
| 218 |
+
haiku_layers.py
|
| 219 |
+
nn_common.py
|
| 220 |
+
```
|
| 221 |
+
|
| 222 |
+
---
|
| 223 |
+
|
| 224 |
+
## Limitations
|
| 225 |
+
|
| 226 |
+
- **Base model**: not instruction-tuned; may ramble or fail at Q&A format
|
| 227 |
+
- **Scale**: ~655M parameters — research / edge prototype, not a frontier model
|
| 228 |
+
- **Dependency**: requires `flash-linear-attention` (KDA); not GGUF / llama.cpp compatible today
|
| 229 |
+
- **Context**: trained at 2,048; longer windows are experimental
|
| 230 |
+
- **Partial epoch**: ~4.4B of 10.13B packed FineWeb-Edu tokens
|
| 231 |
+
|
| 232 |
+
---
|
| 233 |
+
|
| 234 |
+
## Model family
|
| 235 |
+
|
| 236 |
+
| Model | Parameters | Architecture | Stage | Hub |
|
| 237 |
+
|-------|------------|--------------|-------|-----|
|
| 238 |
+
| **Monostich** | ~100M | LLaMA-style | SFT | [`kerzgrr/Monostich`](https://huggingface.co/kerzgrr/Monostich) |
|
| 239 |
+
| **Monostich-2-base** | ~150M | TinyGDN hybrid | Pretrain | [`kerzgrr/Monostich-2-base`](https://huggingface.co/kerzgrr/Monostich-2-base) |
|
| 240 |
+
| **Monostich-2** | ~150M | TinyGDN hybrid | SFT | [`kerzgrr/Monostich-2`](https://huggingface.co/kerzgrr/Monostich-2) |
|
| 241 |
+
| **Couplet-base** | ~268M | TinyGDN hybrid | Pretrain | [`kerzgrr/Couplet-base`](https://huggingface.co/kerzgrr/Couplet-base) |
|
| 242 |
+
| **Couplet** | ~268M | TinyGDN hybrid | SFT | [`kerzgrr/Couplet`](https://huggingface.co/kerzgrr/Couplet) |
|
| 243 |
+
| **Tercet-base** | ~502M | TinyGDN hybrid | Pretrain | [`kerzgrr/Tercet-base`](https://huggingface.co/kerzgrr/Tercet-base) |
|
| 244 |
+
| **Tercet** | ~502M | TinyGDN hybrid | SFT | [`kerzgrr/Tercet`](https://huggingface.co/kerzgrr/Tercet) |
|
| 245 |
+
| **Haiku-base** | ~655M | KDA + gated MLA | Pretrain | *this repo* |
|
| 246 |
+
|
| 247 |
+
---
|
| 248 |
+
|
| 249 |
+
## Citation
|
| 250 |
+
|
| 251 |
+
```bibtex
|
| 252 |
+
@misc{haikubase2026,
|
| 253 |
+
title={Haiku-base: A 655M Hybrid KDA + Gated-MLA Language Model},
|
| 254 |
+
author={kerzgrr},
|
| 255 |
+
year={2026},
|
| 256 |
+
url={https://huggingface.co/kerzgrr/Haiku-base}
|
| 257 |
+
}
|
| 258 |
+
```
|
| 259 |
+
|
| 260 |
+
---
|
| 261 |
+
|
| 262 |
+
## Acknowledgments
|
| 263 |
+
|
| 264 |
+
- [flash-linear-attention](https://github.com/fla-org/flash-linear-attention) (Kimi Delta Attention)
|
| 265 |
+
- [FineWeb-Edu](https://huggingface.co/datasets/HuggingFaceFW/fineweb-edu)
|
| 266 |
+
- Tercet family: [`kerzgrr/Tercet-base`](https://huggingface.co/kerzgrr/Tercet-base)
|
| 267 |
+
- PyTorch SDPA / Hugging Face Hub + tokenizers
|
| 268 |
+
|
| 269 |
+
---
|
| 270 |
+
|
| 271 |
+
<div align="center">
|
| 272 |
+
|
| 273 |
+
*A haiku is three lines — larger than a tercet, still compact.*
|
| 274 |
+
|
| 275 |
+
</div>
|
__pycache__/inference.cpython-312.pyc
ADDED
|
Binary file (22.2 kB). View file
|
|
|
chat_template.jinja
ADDED
|
@@ -0,0 +1,86 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{%- set ns = namespace(xml_tools=none, tools_emitted=false) -%}
|
| 2 |
+
{%- if xml_tools is defined and xml_tools -%}
|
| 3 |
+
{%- set ns.xml_tools = xml_tools -%}
|
| 4 |
+
{%- elif tools is defined and tools -%}
|
| 5 |
+
{%- set ns.xml_tools = tools -%}
|
| 6 |
+
{%- endif -%}
|
| 7 |
+
{%- set tools_preamble = 'You may call one or more functions to assist with the user query.\nYou are provided with function signatures within <tools></tools> XML tags:\n<tools>\n' -%}
|
| 8 |
+
{%- set tools_epilogue = '</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{"name": <function-name>, "arguments": <args-json-object>}\n</tool_call>' -%}
|
| 9 |
+
{%- for message in messages -%}
|
| 10 |
+
{%- if loop.first -%}{{- bos_token -}}{%- endif -%}
|
| 11 |
+
{%- if loop.first and ns.xml_tools and message['role'] != 'system' -%}
|
| 12 |
+
{{- '<|im_start|>system\n' + tools_preamble -}}
|
| 13 |
+
{%- for tool in ns.xml_tools -%}
|
| 14 |
+
{%- if tool is string -%}{{- (tool | replace('<tools>', '') | replace('</tools>', '') | trim) + '\n' -}}
|
| 15 |
+
{%- else -%}{{- tool | tojson + '\n' -}}
|
| 16 |
+
{%- endif -%}
|
| 17 |
+
{%- endfor -%}
|
| 18 |
+
{{- tools_epilogue + '<|im_end|>\n' -}}
|
| 19 |
+
{%- set ns.tools_emitted = true -%}
|
| 20 |
+
{%- endif -%}
|
| 21 |
+
{%- set raw_role = message['role'] -%}
|
| 22 |
+
{%- set role = 'user' if raw_role == 'tool' or raw_role == 'function' else raw_role -%}
|
| 23 |
+
{%- set content_text = namespace(value='') -%}
|
| 24 |
+
{%- if message['content'] is string -%}
|
| 25 |
+
{%- set content_text.value = message['content'] -%}
|
| 26 |
+
{%- elif message['content'] is iterable -%}
|
| 27 |
+
{%- for item in message['content'] -%}
|
| 28 |
+
{%- if item['type'] == 'text' -%}{%- set content_text.value = content_text.value + item['text'] -%}{%- endif -%}
|
| 29 |
+
{%- endfor -%}
|
| 30 |
+
{%- endif -%}
|
| 31 |
+
{%- set is_tool = raw_role == 'tool' or raw_role == 'function' -%}
|
| 32 |
+
{%- set prev_is_tool = loop.previtem is defined and (loop.previtem['role'] == 'tool' or loop.previtem['role'] == 'function') -%}
|
| 33 |
+
{%- set next_is_tool = loop.nextitem is defined and (loop.nextitem['role'] == 'tool' or loop.nextitem['role'] == 'function') -%}
|
| 34 |
+
{%- if raw_role == 'system' and not (content_text.value | trim) and not (ns.xml_tools and not ns.tools_emitted) -%}
|
| 35 |
+
{%- else -%}
|
| 36 |
+
{%- if is_tool and prev_is_tool -%}
|
| 37 |
+
{{- '\n' -}}
|
| 38 |
+
{%- else -%}
|
| 39 |
+
{{- '<|im_start|>' + role + '\n' -}}
|
| 40 |
+
{%- endif -%}
|
| 41 |
+
{%- if raw_role == 'assistant' and enable_thinking is defined -%}
|
| 42 |
+
{{- ('<|think|>\n' if enable_thinking else '<|no_think|>\n') -}}
|
| 43 |
+
{%- endif -%}
|
| 44 |
+
{%- if is_tool -%}
|
| 45 |
+
{{- '<|tool_response|>\n' + (content_text.value | replace('<tool_response>', '') | replace('</tool_response>', '') | trim) -}}
|
| 46 |
+
{%- elif raw_role == 'system' -%}
|
| 47 |
+
{{- content_text.value | replace('/system_override', '') | replace('/no_think', '') | replace('/think', '') | trim -}}
|
| 48 |
+
{%- else -%}
|
| 49 |
+
{{- content_text.value -}}
|
| 50 |
+
{%- endif -%}
|
| 51 |
+
{%- if raw_role == 'system' and ns.xml_tools and not ns.tools_emitted and '<tools>' not in content_text.value -%}
|
| 52 |
+
{%- if content_text.value | trim -%}{{- '\n\n' -}}{%- endif -%}
|
| 53 |
+
{{- tools_preamble -}}
|
| 54 |
+
{%- for tool in ns.xml_tools -%}
|
| 55 |
+
{%- if tool is string -%}{{- (tool | replace('<tools>', '') | replace('</tools>', '') | trim) + '\n' -}}
|
| 56 |
+
{%- else -%}{{- tool | tojson + '\n' -}}
|
| 57 |
+
{%- endif -%}
|
| 58 |
+
{%- endfor -%}
|
| 59 |
+
{{- tools_epilogue -}}
|
| 60 |
+
{%- set ns.tools_emitted = true -%}
|
| 61 |
+
{%- endif -%}
|
| 62 |
+
{%- if raw_role == 'assistant' and message['tool_calls'] is defined and message['tool_calls'] -%}
|
| 63 |
+
{%- if '<tool_call>' not in content_text.value -%}
|
| 64 |
+
{%- for tool_call in message['tool_calls'] -%}
|
| 65 |
+
{%- set fn = tool_call['function'] if tool_call['function'] is defined else tool_call -%}
|
| 66 |
+
{%- if loop.first and not (content_text.value | trim) -%}
|
| 67 |
+
{{- '<tool_call>\n{"name": "' + fn['name'] + '", "arguments": ' -}}
|
| 68 |
+
{%- else -%}
|
| 69 |
+
{{- '\n<tool_call>\n{"name": "' + fn['name'] + '", "arguments": ' -}}
|
| 70 |
+
{%- endif -%}
|
| 71 |
+
{%- if fn['arguments'] is string -%}{{- fn['arguments'] -}}
|
| 72 |
+
{%- else -%}{{- fn['arguments'] | tojson -}}
|
| 73 |
+
{%- endif -%}
|
| 74 |
+
{{- '}\n</tool_call>' -}}
|
| 75 |
+
{%- endfor -%}
|
| 76 |
+
{%- endif -%}
|
| 77 |
+
{%- endif -%}
|
| 78 |
+
{%- if not (is_tool and next_is_tool) -%}
|
| 79 |
+
{{- '<|im_end|>\n' -}}
|
| 80 |
+
{%- endif -%}
|
| 81 |
+
{%- endif -%}
|
| 82 |
+
{%- endfor -%}
|
| 83 |
+
{%- if add_generation_prompt -%}
|
| 84 |
+
{{- '<|im_start|>assistant\n' -}}
|
| 85 |
+
{%- if enable_thinking is defined -%}{{- ('<|think|>\n' if enable_thinking else '<|no_think|>\n') -}}{%- endif -%}
|
| 86 |
+
{%- endif -%}
|
config.json
ADDED
|
@@ -0,0 +1,94 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"allow_negative_eigenvalues": false,
|
| 3 |
+
"architecture": "HaikuForCausalLM",
|
| 4 |
+
"architectures": [
|
| 5 |
+
"HaikuForCausalLM"
|
| 6 |
+
],
|
| 7 |
+
"attention_dropout": 0.0,
|
| 8 |
+
"attention_head_dim": 128,
|
| 9 |
+
"attention_layer_type": "gated_mla",
|
| 10 |
+
"bos_token_id": 0,
|
| 11 |
+
"checkpoint_step": 8400,
|
| 12 |
+
"eos_token_id": 1,
|
| 13 |
+
"full_attention_interval": 4,
|
| 14 |
+
"hidden_size": 1024,
|
| 15 |
+
"initializer_range": 0.02,
|
| 16 |
+
"intermediate_size": 3840,
|
| 17 |
+
"kda_lower_bound": -5.0,
|
| 18 |
+
"kda_safe_gate": true,
|
| 19 |
+
"layer_types": [
|
| 20 |
+
"kda",
|
| 21 |
+
"kda",
|
| 22 |
+
"kda",
|
| 23 |
+
"gated_mla",
|
| 24 |
+
"kda",
|
| 25 |
+
"kda",
|
| 26 |
+
"kda",
|
| 27 |
+
"gated_mla",
|
| 28 |
+
"kda",
|
| 29 |
+
"kda",
|
| 30 |
+
"kda",
|
| 31 |
+
"gated_mla",
|
| 32 |
+
"kda",
|
| 33 |
+
"kda",
|
| 34 |
+
"kda",
|
| 35 |
+
"gated_mla",
|
| 36 |
+
"kda",
|
| 37 |
+
"kda",
|
| 38 |
+
"kda",
|
| 39 |
+
"gated_mla",
|
| 40 |
+
"kda",
|
| 41 |
+
"kda",
|
| 42 |
+
"kda",
|
| 43 |
+
"gated_mla",
|
| 44 |
+
"kda",
|
| 45 |
+
"kda",
|
| 46 |
+
"kda",
|
| 47 |
+
"gated_mla",
|
| 48 |
+
"kda",
|
| 49 |
+
"kda",
|
| 50 |
+
"kda",
|
| 51 |
+
"gated_mla",
|
| 52 |
+
"kda",
|
| 53 |
+
"kda",
|
| 54 |
+
"kda",
|
| 55 |
+
"gated_mla"
|
| 56 |
+
],
|
| 57 |
+
"linear_conv_kernel_dim": 4,
|
| 58 |
+
"linear_expand_v": 1.0,
|
| 59 |
+
"linear_head_dim": 128,
|
| 60 |
+
"linear_layer_type": "kda",
|
| 61 |
+
"linear_num_heads": 8,
|
| 62 |
+
"linear_num_value_heads": 8,
|
| 63 |
+
"max_position_embeddings": 32768,
|
| 64 |
+
"mla_kv_lora_rank": 256,
|
| 65 |
+
"mla_q_lora_rank": 512,
|
| 66 |
+
"mla_qk_nope_head_dim": 128,
|
| 67 |
+
"mla_v_head_dim": 128,
|
| 68 |
+
"mlp_activation": "situ_glu",
|
| 69 |
+
"model_family": "Haiku",
|
| 70 |
+
"model_name": "Haiku-base",
|
| 71 |
+
"model_type": "haiku",
|
| 72 |
+
"mtp_adapter_rank": 128,
|
| 73 |
+
"mtp_loss_weight": 0.2,
|
| 74 |
+
"mtp_num_heads": 2,
|
| 75 |
+
"num_attention_heads": 8,
|
| 76 |
+
"num_hidden_layers": 36,
|
| 77 |
+
"num_key_value_heads": 8,
|
| 78 |
+
"pad_token_id": 2,
|
| 79 |
+
"partial_rotary_factor": 0.5,
|
| 80 |
+
"rms_norm_eps": 1e-06,
|
| 81 |
+
"rope_theta": 1000000.0,
|
| 82 |
+
"shared_layer_indices": [],
|
| 83 |
+
"situ_gate_cap": 4.0,
|
| 84 |
+
"situ_up_cap": 25.0,
|
| 85 |
+
"stage": "pretrain",
|
| 86 |
+
"tie_word_embeddings": true,
|
| 87 |
+
"torch_dtype": "bfloat16",
|
| 88 |
+
"training_sequence_length": 2048,
|
| 89 |
+
"transformers_version": "4.45.0",
|
| 90 |
+
"unk_token_id": 3,
|
| 91 |
+
"use_block_attn_res": true,
|
| 92 |
+
"vocab_size": 65536,
|
| 93 |
+
"weights": "ema"
|
| 94 |
+
}
|
inference.py
ADDED
|
@@ -0,0 +1,501 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
"""Standalone pretrain inference for kerzgrr/Haiku-base.
|
| 3 |
+
|
| 4 |
+
Downloads model assets from the Hub (cached after first run), loads TinyGDN,
|
| 5 |
+
and streams raw text continuations. This is the pretrained base model — not
|
| 6 |
+
instruction-tuned. Chat SFT is not released yet — this repo is the pretrained base only.
|
| 7 |
+
|
| 8 |
+
Examples:
|
| 9 |
+
python inference.py --prompt "The history of computing begins"
|
| 10 |
+
python inference.py
|
| 11 |
+
python inference.py --prompt "Once upon a time" --temperature 0.9 --top-p 0.95
|
| 12 |
+
"""
|
| 13 |
+
|
| 14 |
+
from __future__ import annotations
|
| 15 |
+
|
| 16 |
+
import argparse
|
| 17 |
+
import json
|
| 18 |
+
import os
|
| 19 |
+
import platform
|
| 20 |
+
import shutil
|
| 21 |
+
import subprocess
|
| 22 |
+
import sys
|
| 23 |
+
import time
|
| 24 |
+
import warnings
|
| 25 |
+
from pathlib import Path
|
| 26 |
+
|
| 27 |
+
import torch
|
| 28 |
+
from safetensors.torch import load_file
|
| 29 |
+
from tokenizers import Tokenizer
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def _silence_runtime_warnings() -> None:
|
| 33 |
+
"""Hide known-noisy Triton / SDPA warnings that don't affect results."""
|
| 34 |
+
patterns = (
|
| 35 |
+
r"tl\.make_block_ptr is deprecated",
|
| 36 |
+
r"Memory efficient kernel not used because",
|
| 37 |
+
r"Memory Efficient attention has been runtime disabled",
|
| 38 |
+
r"Flash attention kernel not used because",
|
| 39 |
+
r"Torch was not compiled with flash attention",
|
| 40 |
+
r"cuDNN attention kernel not used because",
|
| 41 |
+
r"cuDNN attention has been runtime disabled",
|
| 42 |
+
)
|
| 43 |
+
for pattern in patterns:
|
| 44 |
+
warnings.filterwarnings("ignore", message=pattern)
|
| 45 |
+
|
| 46 |
+
REPO_ID = "kerzgrr/Haiku-base"
|
| 47 |
+
FLA_COMMIT = "cbb0a72efb55c18ca0ef4f298298317573ad2cb3"
|
| 48 |
+
FLA_REPO = "https://github.com/fla-org/flash-linear-attention.git"
|
| 49 |
+
PATCH_FILES = (
|
| 50 |
+
"fla/__init__.py",
|
| 51 |
+
"fla/ops/__init__.py",
|
| 52 |
+
"fla/layers/__init__.py",
|
| 53 |
+
"fla/ops/simple_gla/__init__.py",
|
| 54 |
+
)
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
def _download(filename: str, local_dir: Path | None) -> Path:
|
| 58 |
+
from huggingface_hub import hf_hub_download
|
| 59 |
+
|
| 60 |
+
return Path(
|
| 61 |
+
hf_hub_download(
|
| 62 |
+
repo_id=REPO_ID,
|
| 63 |
+
filename=filename,
|
| 64 |
+
local_dir=str(local_dir) if local_dir else None,
|
| 65 |
+
)
|
| 66 |
+
)
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
def _run(cmd: list[str], *, cwd: Path | None = None, env: dict | None = None) -> None:
|
| 70 |
+
print("+", " ".join(cmd), flush=True)
|
| 71 |
+
merged = os.environ.copy()
|
| 72 |
+
if env:
|
| 73 |
+
merged.update(env)
|
| 74 |
+
# Avoid Windows cp1252 crashes while reading FLA's README during setup.
|
| 75 |
+
merged.setdefault("PYTHONUTF8", "1")
|
| 76 |
+
merged.setdefault("PYTHONIOENCODING", "utf-8")
|
| 77 |
+
subprocess.check_call(cmd, cwd=str(cwd) if cwd else None, env=merged)
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def _pip_install(*args: str) -> None:
|
| 81 |
+
_run([sys.executable, "-m", "pip", "install", *args])
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
def _fla_importable() -> tuple[bool, str]:
|
| 85 |
+
try:
|
| 86 |
+
from fla.layers.kda import KimiDeltaAttention # noqa: F401
|
| 87 |
+
except Exception as error: # noqa: BLE001
|
| 88 |
+
return False, str(error)
|
| 89 |
+
return True, ""
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
def _cache_root() -> Path:
|
| 93 |
+
override = os.environ.get("MONOSTICH_CACHE")
|
| 94 |
+
if override:
|
| 95 |
+
path = Path(override).expanduser().resolve()
|
| 96 |
+
else:
|
| 97 |
+
path = Path.home() / ".cache" / "haiku-base"
|
| 98 |
+
path.mkdir(parents=True, exist_ok=True)
|
| 99 |
+
return path
|
| 100 |
+
|
| 101 |
+
|
| 102 |
+
def _ensure_git() -> None:
|
| 103 |
+
if shutil.which("git") is None:
|
| 104 |
+
raise RuntimeError(
|
| 105 |
+
"git is required to auto-install flash-linear-attention. "
|
| 106 |
+
"Install Git and ensure it is on PATH."
|
| 107 |
+
)
|
| 108 |
+
|
| 109 |
+
|
| 110 |
+
def _apply_windows_fla_patches(fla_root: Path, local_dir: Path | None) -> None:
|
| 111 |
+
print("Applying Windows FLA import patches from the Hub …", flush=True)
|
| 112 |
+
for relative in PATCH_FILES:
|
| 113 |
+
source = _download(f"windows_fla_patches/{relative}", local_dir)
|
| 114 |
+
target = fla_root / relative
|
| 115 |
+
target.parent.mkdir(parents=True, exist_ok=True)
|
| 116 |
+
shutil.copy2(source, target)
|
| 117 |
+
print(f" patched {relative}", flush=True)
|
| 118 |
+
|
| 119 |
+
|
| 120 |
+
def _install_fla(local_dir: Path | None) -> None:
|
| 121 |
+
print("flash-linear-attention missing/broken — installing automatically …", flush=True)
|
| 122 |
+
_pip_install("einops", "numpy")
|
| 123 |
+
is_windows = platform.system() == "Windows"
|
| 124 |
+
|
| 125 |
+
if not is_windows:
|
| 126 |
+
_pip_install(
|
| 127 |
+
"--no-deps",
|
| 128 |
+
f"git+{FLA_REPO}@{FLA_COMMIT}",
|
| 129 |
+
)
|
| 130 |
+
return
|
| 131 |
+
|
| 132 |
+
# Windows: editable clone + Hub patches (stock FLA + Triton 3.7 breaks on import).
|
| 133 |
+
_ensure_git()
|
| 134 |
+
fla_root = _cache_root() / "flash-linear-attention"
|
| 135 |
+
if (fla_root / ".git").is_dir():
|
| 136 |
+
_run(["git", "fetch", "--depth", "1", "origin", FLA_COMMIT], cwd=fla_root)
|
| 137 |
+
_run(["git", "checkout", "--force", FLA_COMMIT], cwd=fla_root)
|
| 138 |
+
else:
|
| 139 |
+
if fla_root.exists():
|
| 140 |
+
shutil.rmtree(fla_root)
|
| 141 |
+
_run(
|
| 142 |
+
[
|
| 143 |
+
"git",
|
| 144 |
+
"clone",
|
| 145 |
+
"--filter=blob:none",
|
| 146 |
+
FLA_REPO,
|
| 147 |
+
str(fla_root),
|
| 148 |
+
]
|
| 149 |
+
)
|
| 150 |
+
_run(["git", "fetch", "--depth", "1", "origin", FLA_COMMIT], cwd=fla_root)
|
| 151 |
+
_run(["git", "checkout", "--force", FLA_COMMIT], cwd=fla_root)
|
| 152 |
+
|
| 153 |
+
_apply_windows_fla_patches(fla_root, local_dir)
|
| 154 |
+
_pip_install("--no-build-isolation", "--no-deps", "-e", str(fla_root))
|
| 155 |
+
|
| 156 |
+
|
| 157 |
+
def _ensure_fla(local_dir: Path | None) -> None:
|
| 158 |
+
ok, error = _fla_importable()
|
| 159 |
+
if ok:
|
| 160 |
+
return
|
| 161 |
+
print(f"FLA not ready ({error})", flush=True)
|
| 162 |
+
try:
|
| 163 |
+
_install_fla(local_dir)
|
| 164 |
+
except Exception as install_error: # noqa: BLE001
|
| 165 |
+
raise RuntimeError(
|
| 166 |
+
"Automatic flash-linear-attention install failed.\n"
|
| 167 |
+
f"Original import error: {error}\n"
|
| 168 |
+
f"Install error: {install_error}"
|
| 169 |
+
) from install_error
|
| 170 |
+
|
| 171 |
+
# Drop cached failed imports so the newly installed package is picked up.
|
| 172 |
+
for name in list(sys.modules):
|
| 173 |
+
if name == "fla" or name.startswith("fla."):
|
| 174 |
+
del sys.modules[name]
|
| 175 |
+
|
| 176 |
+
ok, error = _fla_importable()
|
| 177 |
+
if not ok:
|
| 178 |
+
raise RuntimeError(
|
| 179 |
+
"flash-linear-attention installed but still failed to import "
|
| 180 |
+
f"KimiDeltaAttention: {error}"
|
| 181 |
+
)
|
| 182 |
+
print("flash-linear-attention ready.", flush=True)
|
| 183 |
+
|
| 184 |
+
|
| 185 |
+
def _ensure_tiny_gdn(local_dir: Path | None) -> Path:
|
| 186 |
+
"""Return a directory that contains the tiny_gdn package on sys.path."""
|
| 187 |
+
# Prefer files next to this script (local clone / Hub snapshot).
|
| 188 |
+
here = Path(__file__).resolve().parent
|
| 189 |
+
if (here / "tiny_gdn" / "__init__.py").is_file():
|
| 190 |
+
return here
|
| 191 |
+
if local_dir and (local_dir / "tiny_gdn" / "__init__.py").is_file():
|
| 192 |
+
return local_dir
|
| 193 |
+
|
| 194 |
+
# Pull package modules from the Hub into the cache.
|
| 195 |
+
for name in (
|
| 196 |
+
"tiny_gdn/__init__.py",
|
| 197 |
+
"tiny_gdn/config.py",
|
| 198 |
+
"tiny_gdn/model.py",
|
| 199 |
+
"tiny_gdn/haiku_layers.py",
|
| 200 |
+
"tiny_gdn/nn_common.py",
|
| 201 |
+
):
|
| 202 |
+
_download(name, local_dir)
|
| 203 |
+
# hf_hub_download with local_dir=None caches under hub/; resolve via download of init.
|
| 204 |
+
init_path = _download("tiny_gdn/__init__.py", local_dir)
|
| 205 |
+
return init_path.parent.parent
|
| 206 |
+
|
| 207 |
+
|
| 208 |
+
def _sample(
|
| 209 |
+
logits: torch.Tensor,
|
| 210 |
+
*,
|
| 211 |
+
temperature: float,
|
| 212 |
+
top_p: float,
|
| 213 |
+
top_k: int,
|
| 214 |
+
generator: torch.Generator,
|
| 215 |
+
) -> int:
|
| 216 |
+
logits = logits.float()
|
| 217 |
+
if temperature <= 1e-5:
|
| 218 |
+
return int(torch.argmax(logits).item())
|
| 219 |
+
logits = logits / temperature
|
| 220 |
+
if 0 < top_k < logits.shape[-1]:
|
| 221 |
+
threshold = torch.topk(logits, top_k).values[-1]
|
| 222 |
+
logits = logits.masked_fill(logits < threshold, -torch.inf)
|
| 223 |
+
if top_p < 1.0:
|
| 224 |
+
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
|
| 225 |
+
probs = torch.softmax(sorted_logits, dim=-1)
|
| 226 |
+
remove = torch.cumsum(probs, dim=-1) > top_p
|
| 227 |
+
remove[1:] = remove[:-1].clone()
|
| 228 |
+
remove[0] = False
|
| 229 |
+
sorted_logits = sorted_logits.masked_fill(remove, -torch.inf)
|
| 230 |
+
logits = torch.full_like(logits, -torch.inf)
|
| 231 |
+
logits.scatter_(0, sorted_indices, sorted_logits)
|
| 232 |
+
probs = torch.softmax(logits, dim=-1)
|
| 233 |
+
return int(torch.multinomial(probs, 1, generator=generator).item())
|
| 234 |
+
|
| 235 |
+
|
| 236 |
+
def _apply_repetition_penalty(
|
| 237 |
+
logits: torch.Tensor,
|
| 238 |
+
token_ids: list[int],
|
| 239 |
+
penalty: float,
|
| 240 |
+
window: int,
|
| 241 |
+
) -> torch.Tensor:
|
| 242 |
+
if penalty == 1.0 or not token_ids:
|
| 243 |
+
return logits
|
| 244 |
+
recent = token_ids[-window:] if window > 0 else token_ids
|
| 245 |
+
unique = torch.tensor(list(set(recent)), dtype=torch.long, device=logits.device)
|
| 246 |
+
score = logits[unique]
|
| 247 |
+
logits[unique] = torch.where(score > 0, score / penalty, score * penalty)
|
| 248 |
+
return logits
|
| 249 |
+
|
| 250 |
+
|
| 251 |
+
@torch.inference_mode()
|
| 252 |
+
def generate(
|
| 253 |
+
model,
|
| 254 |
+
tokenizer: Tokenizer,
|
| 255 |
+
prompt_ids: list[int],
|
| 256 |
+
*,
|
| 257 |
+
max_new_tokens: int,
|
| 258 |
+
context_length: int,
|
| 259 |
+
temperature: float,
|
| 260 |
+
top_p: float,
|
| 261 |
+
top_k: int,
|
| 262 |
+
repetition_penalty: float,
|
| 263 |
+
repetition_window: int,
|
| 264 |
+
seed: int,
|
| 265 |
+
stream: bool,
|
| 266 |
+
device: torch.device,
|
| 267 |
+
) -> tuple[str, int, str]:
|
| 268 |
+
eos_id = int(model.config.eos_token_id)
|
| 269 |
+
token_ids = list(prompt_ids[-context_length:])
|
| 270 |
+
generated: list[int] = []
|
| 271 |
+
decoded = ""
|
| 272 |
+
stop_reason = "max_new_tokens"
|
| 273 |
+
generator = torch.Generator(device=device)
|
| 274 |
+
generator.manual_seed(seed)
|
| 275 |
+
started = time.perf_counter()
|
| 276 |
+
|
| 277 |
+
for _ in range(max_new_tokens):
|
| 278 |
+
context = token_ids[-context_length:]
|
| 279 |
+
input_ids = torch.tensor([context], dtype=torch.long, device=device)
|
| 280 |
+
output = model(input_ids, return_logits=True, logits_to_keep=1)
|
| 281 |
+
if output.logits is None:
|
| 282 |
+
raise RuntimeError("Model returned no logits")
|
| 283 |
+
next_logits = _apply_repetition_penalty(
|
| 284 |
+
output.logits[0, -1],
|
| 285 |
+
token_ids,
|
| 286 |
+
repetition_penalty,
|
| 287 |
+
repetition_window,
|
| 288 |
+
)
|
| 289 |
+
next_id = _sample(
|
| 290 |
+
next_logits,
|
| 291 |
+
temperature=temperature,
|
| 292 |
+
top_p=top_p,
|
| 293 |
+
top_k=top_k,
|
| 294 |
+
generator=generator,
|
| 295 |
+
)
|
| 296 |
+
if next_id == eos_id:
|
| 297 |
+
stop_reason = "eos"
|
| 298 |
+
break
|
| 299 |
+
token_ids.append(next_id)
|
| 300 |
+
generated.append(next_id)
|
| 301 |
+
current = tokenizer.decode(generated, skip_special_tokens=True)
|
| 302 |
+
delta = (
|
| 303 |
+
current[len(decoded) :]
|
| 304 |
+
if current.startswith(decoded)
|
| 305 |
+
else current
|
| 306 |
+
)
|
| 307 |
+
decoded = current
|
| 308 |
+
if stream and delta:
|
| 309 |
+
print(delta, end="", flush=True)
|
| 310 |
+
|
| 311 |
+
if stream:
|
| 312 |
+
print(flush=True)
|
| 313 |
+
elapsed = time.perf_counter() - started
|
| 314 |
+
tps = len(generated) / max(elapsed, 1e-9)
|
| 315 |
+
if stream:
|
| 316 |
+
print(
|
| 317 |
+
f"[done] tokens={len(generated)} stop={stop_reason} "
|
| 318 |
+
f"{tps:.1f} tok/s",
|
| 319 |
+
file=sys.stderr,
|
| 320 |
+
flush=True,
|
| 321 |
+
)
|
| 322 |
+
return decoded, len(generated), stop_reason
|
| 323 |
+
|
| 324 |
+
|
| 325 |
+
def parse_args() -> argparse.Namespace:
|
| 326 |
+
parser = argparse.ArgumentParser(
|
| 327 |
+
description="Haiku-base pretrained text continuation"
|
| 328 |
+
)
|
| 329 |
+
parser.add_argument(
|
| 330 |
+
"--prompt",
|
| 331 |
+
default=None,
|
| 332 |
+
help="Single prompt to continue (omit for interactive REPL)",
|
| 333 |
+
)
|
| 334 |
+
parser.add_argument("--max-new-tokens", type=int, default=256)
|
| 335 |
+
parser.add_argument("--temperature", type=float, default=0.8)
|
| 336 |
+
parser.add_argument("--top-p", type=float, default=0.95)
|
| 337 |
+
parser.add_argument("--top-k", type=int, default=50)
|
| 338 |
+
parser.add_argument("--repetition-penalty", type=float, default=1.08)
|
| 339 |
+
parser.add_argument("--repetition-window", type=int, default=256)
|
| 340 |
+
parser.add_argument("--context-length", type=int, default=2048)
|
| 341 |
+
parser.add_argument("--seed", type=int, default=42)
|
| 342 |
+
parser.add_argument(
|
| 343 |
+
"--device",
|
| 344 |
+
default="cuda" if torch.cuda.is_available() else "cpu",
|
| 345 |
+
choices=["cuda", "cpu"],
|
| 346 |
+
)
|
| 347 |
+
parser.add_argument(
|
| 348 |
+
"--no-stream",
|
| 349 |
+
action="store_true",
|
| 350 |
+
help="Disable token streaming (print full completion at once)",
|
| 351 |
+
)
|
| 352 |
+
parser.add_argument(
|
| 353 |
+
"--no-bos",
|
| 354 |
+
action="store_true",
|
| 355 |
+
help="Do not prepend <|begin_of_text|> to the prompt",
|
| 356 |
+
)
|
| 357 |
+
parser.add_argument(
|
| 358 |
+
"--local-dir",
|
| 359 |
+
default=None,
|
| 360 |
+
help="Optional local snapshot directory (skips re-download when populated)",
|
| 361 |
+
)
|
| 362 |
+
parser.add_argument(
|
| 363 |
+
"--repo-id",
|
| 364 |
+
default=REPO_ID,
|
| 365 |
+
help=f"Hub repo id (default: {REPO_ID})",
|
| 366 |
+
)
|
| 367 |
+
return parser.parse_args()
|
| 368 |
+
|
| 369 |
+
|
| 370 |
+
def main() -> int:
|
| 371 |
+
_silence_runtime_warnings()
|
| 372 |
+
args = parse_args()
|
| 373 |
+
global REPO_ID
|
| 374 |
+
REPO_ID = args.repo_id
|
| 375 |
+
local_dir = Path(args.local_dir).resolve() if args.local_dir else None
|
| 376 |
+
|
| 377 |
+
print(f"Loading Haiku-base from huggingface.co/{REPO_ID} …", flush=True)
|
| 378 |
+
try:
|
| 379 |
+
package_root = _ensure_tiny_gdn(local_dir)
|
| 380 |
+
except Exception as error: # noqa: BLE001
|
| 381 |
+
print(f"Failed to resolve tiny_gdn package: {error}", file=sys.stderr)
|
| 382 |
+
print(
|
| 383 |
+
"Make sure huggingface_hub is installed and you can reach the Hub.",
|
| 384 |
+
file=sys.stderr,
|
| 385 |
+
)
|
| 386 |
+
return 1
|
| 387 |
+
|
| 388 |
+
if str(package_root) not in sys.path:
|
| 389 |
+
sys.path.insert(0, str(package_root))
|
| 390 |
+
|
| 391 |
+
try:
|
| 392 |
+
_ensure_fla(local_dir)
|
| 393 |
+
except Exception as error: # noqa: BLE001
|
| 394 |
+
print(str(error), file=sys.stderr)
|
| 395 |
+
return 1
|
| 396 |
+
|
| 397 |
+
try:
|
| 398 |
+
from tiny_gdn import TinyGDNConfig, TinyGDNForCausalLM
|
| 399 |
+
except ImportError as error:
|
| 400 |
+
print(f"Could not import tiny_gdn: {error}", file=sys.stderr)
|
| 401 |
+
return 1
|
| 402 |
+
|
| 403 |
+
weights_path = _download("model.safetensors", local_dir)
|
| 404 |
+
tok_path = _download("tokenizer.json", local_dir)
|
| 405 |
+
cfg_path = _download("config.json", local_dir)
|
| 406 |
+
|
| 407 |
+
raw = json.loads(cfg_path.read_text(encoding="utf-8"))
|
| 408 |
+
# Config.from_json rejects unknown HF-only keys — filter to dataclass fields.
|
| 409 |
+
from dataclasses import fields
|
| 410 |
+
|
| 411 |
+
allowed = {item.name for item in fields(TinyGDNConfig)}
|
| 412 |
+
payload = {key: value for key, value in raw.items() if key in allowed}
|
| 413 |
+
if "shared_layer_indices" in payload:
|
| 414 |
+
payload["shared_layer_indices"] = tuple(payload["shared_layer_indices"])
|
| 415 |
+
config = TinyGDNConfig(**payload)
|
| 416 |
+
|
| 417 |
+
device = torch.device(args.device)
|
| 418 |
+
if device.type == "cuda" and not torch.cuda.is_available():
|
| 419 |
+
print("CUDA requested but unavailable; falling back to CPU.", flush=True)
|
| 420 |
+
device = torch.device("cpu")
|
| 421 |
+
dtype = torch.bfloat16 if device.type == "cuda" else torch.float32
|
| 422 |
+
|
| 423 |
+
print(
|
| 424 |
+
f"Building Haiku ({config.num_hidden_layers}L / {config.hidden_size}d) "
|
| 425 |
+
f"on {device} …",
|
| 426 |
+
flush=True,
|
| 427 |
+
)
|
| 428 |
+
model = TinyGDNForCausalLM(config)
|
| 429 |
+
state = load_file(str(weights_path), device="cpu")
|
| 430 |
+
model.load_state_dict(state, strict=True)
|
| 431 |
+
del state
|
| 432 |
+
model = model.to(device=device, dtype=dtype)
|
| 433 |
+
model.eval()
|
| 434 |
+
model.requires_grad_(False)
|
| 435 |
+
|
| 436 |
+
tokenizer = Tokenizer.from_file(str(tok_path))
|
| 437 |
+
bos_id = config.bos_token_id
|
| 438 |
+
context_length = min(args.context_length, config.max_position_embeddings)
|
| 439 |
+
stream = not args.no_stream
|
| 440 |
+
|
| 441 |
+
def encode_prompt(text: str) -> list[int]:
|
| 442 |
+
ids = tokenizer.encode(text, add_special_tokens=False).ids
|
| 443 |
+
if not args.no_bos and (not ids or ids[0] != bos_id):
|
| 444 |
+
ids = [bos_id] + ids
|
| 445 |
+
return ids
|
| 446 |
+
|
| 447 |
+
def run_once(prompt: str) -> None:
|
| 448 |
+
prompt_ids = encode_prompt(prompt)
|
| 449 |
+
if stream:
|
| 450 |
+
print(prompt, end="", flush=True)
|
| 451 |
+
text, _, _ = generate(
|
| 452 |
+
model,
|
| 453 |
+
tokenizer,
|
| 454 |
+
prompt_ids,
|
| 455 |
+
max_new_tokens=args.max_new_tokens,
|
| 456 |
+
context_length=context_length,
|
| 457 |
+
temperature=args.temperature,
|
| 458 |
+
top_p=args.top_p,
|
| 459 |
+
top_k=args.top_k,
|
| 460 |
+
repetition_penalty=args.repetition_penalty,
|
| 461 |
+
repetition_window=args.repetition_window,
|
| 462 |
+
seed=args.seed,
|
| 463 |
+
stream=stream,
|
| 464 |
+
device=device,
|
| 465 |
+
)
|
| 466 |
+
if not stream:
|
| 467 |
+
print(prompt + text)
|
| 468 |
+
|
| 469 |
+
if args.prompt is not None:
|
| 470 |
+
run_once(args.prompt)
|
| 471 |
+
return 0
|
| 472 |
+
|
| 473 |
+
print(
|
| 474 |
+
"Interactive pretrain continuation. Type a prompt and press Enter.\n"
|
| 475 |
+
"Commands: /exit /quit /reset (clears nothing persistent; just a noop marker)\n"
|
| 476 |
+
"Note: this is the BASE model — raw continuation, not chat.",
|
| 477 |
+
flush=True,
|
| 478 |
+
)
|
| 479 |
+
while True:
|
| 480 |
+
try:
|
| 481 |
+
user_input = input("prompt> ")
|
| 482 |
+
except (EOFError, KeyboardInterrupt):
|
| 483 |
+
print()
|
| 484 |
+
break
|
| 485 |
+
text = user_input.strip()
|
| 486 |
+
if not text:
|
| 487 |
+
continue
|
| 488 |
+
if text.lower() in {"/exit", "/quit"}:
|
| 489 |
+
break
|
| 490 |
+
if text.lower() == "/reset":
|
| 491 |
+
print("(history not kept in pretrain mode)", flush=True)
|
| 492 |
+
continue
|
| 493 |
+
if stream:
|
| 494 |
+
print("completion> ", end="", flush=True)
|
| 495 |
+
run_once(text)
|
| 496 |
+
print(flush=True)
|
| 497 |
+
return 0
|
| 498 |
+
|
| 499 |
+
|
| 500 |
+
if __name__ == "__main__":
|
| 501 |
+
raise SystemExit(main())
|
merges.txt
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:64cc003e3d5e29aad7f42a8da43b094b1d77dbe76d5d1d4220df518dd364d265
|
| 3 |
+
size 1311667272
|
requirements.txt
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
torch>=2.4.0
|
| 2 |
+
safetensors>=0.4.0
|
| 3 |
+
tokenizers>=0.20.0
|
| 4 |
+
huggingface_hub>=0.26.0
|
| 5 |
+
einops>=0.8.0
|
| 6 |
+
numpy>=1.26.0
|
| 7 |
+
|
| 8 |
+
# Install separately (pinned commit used for Haiku):
|
| 9 |
+
# pip install --no-deps "git+https://github.com/fla-org/flash-linear-attention.git@cbb0a72efb55c18ca0ef4f298298317573ad2cb3"
|
special_token_ids.json
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"bos_token": "<|begin_of_text|>",
|
| 3 |
+
"eos_token": "<|end_of_text|>",
|
| 4 |
+
"pad_token": "<|padding|>",
|
| 5 |
+
"unk_token": "<|unknown|>",
|
| 6 |
+
"bos_token_id": 0,
|
| 7 |
+
"eos_token_id": 1,
|
| 8 |
+
"pad_token_id": 2,
|
| 9 |
+
"unk_token_id": 3
|
| 10 |
+
}
|
special_tokens_map.json
ADDED
|
@@ -0,0 +1,68 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"bos_token": "<|begin_of_text|>",
|
| 3 |
+
"eos_token": "<|end_of_text|>",
|
| 4 |
+
"pad_token": "<|padding|>",
|
| 5 |
+
"unk_token": "<|unknown|>",
|
| 6 |
+
"additional_special_tokens": [
|
| 7 |
+
"<|im_start|>",
|
| 8 |
+
"<|im_end|>",
|
| 9 |
+
"<|tool_call|>",
|
| 10 |
+
"<|tool_response|>",
|
| 11 |
+
"<think>",
|
| 12 |
+
"</think>",
|
| 13 |
+
"<|fim_prefix|>",
|
| 14 |
+
"<|fim_middle|>",
|
| 15 |
+
"<|fim_suffix|>",
|
| 16 |
+
"<|fim_pad|>",
|
| 17 |
+
"<|no_think|>",
|
| 18 |
+
"<|think|>",
|
| 19 |
+
"<|reserved_002|>",
|
| 20 |
+
"<|reserved_003|>",
|
| 21 |
+
"<|reserved_004|>",
|
| 22 |
+
"<|reserved_005|>",
|
| 23 |
+
"<|reserved_006|>",
|
| 24 |
+
"<|reserved_007|>",
|
| 25 |
+
"<|reserved_008|>",
|
| 26 |
+
"<|reserved_009|>",
|
| 27 |
+
"<|reserved_010|>",
|
| 28 |
+
"<|reserved_011|>",
|
| 29 |
+
"<|reserved_012|>",
|
| 30 |
+
"<|reserved_013|>",
|
| 31 |
+
"<|reserved_014|>",
|
| 32 |
+
"<|reserved_015|>",
|
| 33 |
+
"<|reserved_016|>",
|
| 34 |
+
"<|reserved_017|>",
|
| 35 |
+
"<|reserved_018|>",
|
| 36 |
+
"<|reserved_019|>",
|
| 37 |
+
"<|reserved_020|>",
|
| 38 |
+
"<|reserved_021|>",
|
| 39 |
+
"<|reserved_022|>",
|
| 40 |
+
"<|reserved_023|>",
|
| 41 |
+
"<|reserved_024|>",
|
| 42 |
+
"<|reserved_025|>",
|
| 43 |
+
"<|reserved_026|>",
|
| 44 |
+
"<|reserved_027|>",
|
| 45 |
+
"<|reserved_028|>",
|
| 46 |
+
"<|reserved_029|>",
|
| 47 |
+
"<|reserved_030|>",
|
| 48 |
+
"<|reserved_031|>",
|
| 49 |
+
"<|reserved_032|>",
|
| 50 |
+
"<|reserved_033|>",
|
| 51 |
+
"<|reserved_034|>",
|
| 52 |
+
"<|reserved_035|>",
|
| 53 |
+
"<|reserved_036|>",
|
| 54 |
+
"<|reserved_037|>",
|
| 55 |
+
"<|reserved_038|>",
|
| 56 |
+
"<|reserved_039|>",
|
| 57 |
+
"<|reserved_040|>",
|
| 58 |
+
"<|reserved_041|>",
|
| 59 |
+
"<|reserved_042|>",
|
| 60 |
+
"<|reserved_043|>",
|
| 61 |
+
"<|reserved_044|>",
|
| 62 |
+
"<|reserved_045|>",
|
| 63 |
+
"<|reserved_046|>",
|
| 64 |
+
"<|reserved_047|>",
|
| 65 |
+
"<|reserved_048|>",
|
| 66 |
+
"<|reserved_049|>"
|
| 67 |
+
]
|
| 68 |
+
}
|
tiny_gdn/__init__.py
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from tiny_gdn.config import TinyGDNConfig, haiku_config
|
| 2 |
+
from tiny_gdn.model import TinyGDNForCausalLM, TinyGDNOutput
|
| 3 |
+
|
| 4 |
+
__all__ = [
|
| 5 |
+
"TinyGDNConfig",
|
| 6 |
+
"TinyGDNForCausalLM",
|
| 7 |
+
"TinyGDNOutput",
|
| 8 |
+
"haiku_config",
|
| 9 |
+
]
|
tiny_gdn/config.py
ADDED
|
@@ -0,0 +1,193 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import json
|
| 4 |
+
from dataclasses import asdict, dataclass, fields
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
from typing import Any
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
@dataclass(frozen=True)
|
| 10 |
+
class TinyGDNConfig:
|
| 11 |
+
architecture: str = "TinyGDNForCausalLM"
|
| 12 |
+
model_type: str = "tiny_gdn"
|
| 13 |
+
|
| 14 |
+
vocab_size: int = 49_152
|
| 15 |
+
# Deep-thin sizing is deliberate: controlled sub-billion studies find
|
| 16 |
+
# depth materially more valuable than width around the 125M-150M scale.
|
| 17 |
+
hidden_size: int = 512
|
| 18 |
+
intermediate_size: int = 1_472
|
| 19 |
+
num_hidden_layers: int = 32
|
| 20 |
+
|
| 21 |
+
num_attention_heads: int = 4
|
| 22 |
+
num_key_value_heads: int = 1
|
| 23 |
+
attention_head_dim: int = 128
|
| 24 |
+
full_attention_interval: int = 4
|
| 25 |
+
attention_dropout: float = 0.0
|
| 26 |
+
partial_rotary_factor: float = 0.5
|
| 27 |
+
rope_theta: float = 1_000_000.0
|
| 28 |
+
|
| 29 |
+
linear_num_heads: int = 4
|
| 30 |
+
linear_num_value_heads: int = 4
|
| 31 |
+
linear_head_dim: int = 128
|
| 32 |
+
linear_expand_v: float = 1.0
|
| 33 |
+
linear_conv_kernel_dim: int = 4
|
| 34 |
+
allow_negative_eigenvalues: bool = False
|
| 35 |
+
linear_layer_type: str = "gdn2"
|
| 36 |
+
attention_layer_type: str = "full_attention"
|
| 37 |
+
mlp_activation: str = "swiglu"
|
| 38 |
+
use_block_attn_res: bool = False
|
| 39 |
+
kda_safe_gate: bool = True
|
| 40 |
+
kda_lower_bound: float = -5.0
|
| 41 |
+
mla_q_lora_rank: int | None = 512
|
| 42 |
+
mla_kv_lora_rank: int = 256
|
| 43 |
+
mla_qk_nope_head_dim: int = 128
|
| 44 |
+
mla_v_head_dim: int = 128
|
| 45 |
+
situ_gate_cap: float = 4.0
|
| 46 |
+
situ_up_cap: float = 25.0
|
| 47 |
+
|
| 48 |
+
max_position_embeddings: int = 32_768
|
| 49 |
+
training_sequence_length: int = 2_048
|
| 50 |
+
rms_norm_eps: float = 1e-6
|
| 51 |
+
initializer_range: float = 0.02
|
| 52 |
+
tie_word_embeddings: bool = True
|
| 53 |
+
shared_layer_indices: tuple[int, ...] = ()
|
| 54 |
+
|
| 55 |
+
# MTP is an opt-in ablation at this scale; static MTP is not assumed to
|
| 56 |
+
# improve a 150M model without a controlled pilot.
|
| 57 |
+
mtp_num_heads: int = 0
|
| 58 |
+
mtp_adapter_rank: int = 128
|
| 59 |
+
mtp_loss_weight: float = 0.0
|
| 60 |
+
|
| 61 |
+
bos_token_id: int = 0
|
| 62 |
+
eos_token_id: int = 1
|
| 63 |
+
pad_token_id: int = 2
|
| 64 |
+
unk_token_id: int = 3
|
| 65 |
+
|
| 66 |
+
def __post_init__(self) -> None:
|
| 67 |
+
if self.vocab_size <= 0 or self.vocab_size > 65_536:
|
| 68 |
+
raise ValueError("vocab_size must fit the uint16 token dataset")
|
| 69 |
+
if self.hidden_size != self.num_attention_heads * self.attention_head_dim:
|
| 70 |
+
raise ValueError("hidden_size must equal num_attention_heads * attention_head_dim")
|
| 71 |
+
if self.hidden_size != self.linear_num_heads * self.linear_head_dim:
|
| 72 |
+
raise ValueError("hidden_size must equal linear_num_heads * linear_head_dim")
|
| 73 |
+
if self.linear_num_value_heads < self.linear_num_heads:
|
| 74 |
+
raise ValueError("linear_num_value_heads must be at least linear_num_heads")
|
| 75 |
+
if self.linear_num_value_heads % self.linear_num_heads != 0:
|
| 76 |
+
raise ValueError("linear_num_value_heads must be divisible by linear_num_heads")
|
| 77 |
+
if self.num_attention_heads % self.num_key_value_heads != 0:
|
| 78 |
+
raise ValueError("num_attention_heads must be divisible by num_key_value_heads")
|
| 79 |
+
if self.num_hidden_layers % self.full_attention_interval != 0:
|
| 80 |
+
raise ValueError("num_hidden_layers must be divisible by full_attention_interval")
|
| 81 |
+
if not 0.0 < self.partial_rotary_factor <= 1.0:
|
| 82 |
+
raise ValueError("partial_rotary_factor must be in (0, 1]")
|
| 83 |
+
rotary_dim = int(self.attention_head_dim * self.partial_rotary_factor)
|
| 84 |
+
if rotary_dim <= 0 or rotary_dim % 2:
|
| 85 |
+
raise ValueError("The partial rotary dimension must be positive and even")
|
| 86 |
+
if self.training_sequence_length > self.max_position_embeddings:
|
| 87 |
+
raise ValueError("training_sequence_length exceeds max_position_embeddings")
|
| 88 |
+
if len(set(self.shared_layer_indices)) != len(self.shared_layer_indices):
|
| 89 |
+
raise ValueError("shared_layer_indices must be unique")
|
| 90 |
+
if any(
|
| 91 |
+
index < 0 or index >= self.num_hidden_layers
|
| 92 |
+
for index in self.shared_layer_indices
|
| 93 |
+
):
|
| 94 |
+
raise ValueError("shared_layer_indices contains an invalid layer")
|
| 95 |
+
if self.mtp_num_heads < 0:
|
| 96 |
+
raise ValueError("mtp_num_heads cannot be negative")
|
| 97 |
+
if self.mtp_num_heads and self.mtp_adapter_rank <= 0:
|
| 98 |
+
raise ValueError("mtp_adapter_rank must be positive when MTP is enabled")
|
| 99 |
+
if not 0.0 <= self.mtp_loss_weight <= 1.0:
|
| 100 |
+
raise ValueError("mtp_loss_weight must be between zero and one")
|
| 101 |
+
for token_id in (
|
| 102 |
+
self.bos_token_id,
|
| 103 |
+
self.eos_token_id,
|
| 104 |
+
self.pad_token_id,
|
| 105 |
+
self.unk_token_id,
|
| 106 |
+
):
|
| 107 |
+
if not 0 <= token_id < self.vocab_size:
|
| 108 |
+
raise ValueError(f"Special token ID {token_id} is outside the vocabulary")
|
| 109 |
+
if self.linear_layer_type not in {"gdn2", "kda"}:
|
| 110 |
+
raise ValueError("linear_layer_type must be 'gdn2' or 'kda'")
|
| 111 |
+
if self.attention_layer_type not in {"full_attention", "gated_mla"}:
|
| 112 |
+
raise ValueError("attention_layer_type must be 'full_attention' or 'gated_mla'")
|
| 113 |
+
if self.mlp_activation not in {"swiglu", "situ_glu"}:
|
| 114 |
+
raise ValueError("mlp_activation must be 'swiglu' or 'situ_glu'")
|
| 115 |
+
if self.mla_kv_lora_rank <= 0:
|
| 116 |
+
raise ValueError("mla_kv_lora_rank must be positive")
|
| 117 |
+
if self.mla_q_lora_rank is not None and self.mla_q_lora_rank <= 0:
|
| 118 |
+
raise ValueError("mla_q_lora_rank must be positive or None")
|
| 119 |
+
if self.mla_qk_nope_head_dim <= 0 or self.mla_v_head_dim <= 0:
|
| 120 |
+
raise ValueError("MLA head dimensions must be positive")
|
| 121 |
+
if self.kda_lower_bound >= 0.0:
|
| 122 |
+
raise ValueError("kda_lower_bound must be negative log-decay")
|
| 123 |
+
if self.situ_gate_cap <= 0.0 or self.situ_up_cap <= 0.0:
|
| 124 |
+
raise ValueError("SiTU caps must be positive")
|
| 125 |
+
|
| 126 |
+
@property
|
| 127 |
+
def layer_types(self) -> tuple[str, ...]:
|
| 128 |
+
return tuple(
|
| 129 |
+
self.attention_layer_type
|
| 130 |
+
if (index + 1) % self.full_attention_interval == 0
|
| 131 |
+
else self.linear_layer_type
|
| 132 |
+
for index in range(self.num_hidden_layers)
|
| 133 |
+
)
|
| 134 |
+
|
| 135 |
+
@property
|
| 136 |
+
def rotary_dim(self) -> int:
|
| 137 |
+
return int(self.attention_head_dim * self.partial_rotary_factor)
|
| 138 |
+
|
| 139 |
+
@property
|
| 140 |
+
def effective_num_layers(self) -> int:
|
| 141 |
+
return self.num_hidden_layers + len(self.shared_layer_indices)
|
| 142 |
+
|
| 143 |
+
def to_dict(self) -> dict[str, Any]:
|
| 144 |
+
payload = asdict(self)
|
| 145 |
+
payload["layer_types"] = list(self.layer_types)
|
| 146 |
+
return payload
|
| 147 |
+
|
| 148 |
+
def save_json(self, path: Path) -> None:
|
| 149 |
+
path.parent.mkdir(parents=True, exist_ok=True)
|
| 150 |
+
path.write_text(
|
| 151 |
+
json.dumps(self.to_dict(), indent=2, sort_keys=True) + "\n",
|
| 152 |
+
encoding="utf-8",
|
| 153 |
+
)
|
| 154 |
+
|
| 155 |
+
@classmethod
|
| 156 |
+
def from_json(cls, path: Path) -> TinyGDNConfig:
|
| 157 |
+
payload = json.loads(path.read_text(encoding="utf-8"))
|
| 158 |
+
payload.pop("layer_types", None)
|
| 159 |
+
if "shared_layer_indices" in payload:
|
| 160 |
+
payload["shared_layer_indices"] = tuple(payload["shared_layer_indices"])
|
| 161 |
+
allowed = {item.name for item in fields(cls)}
|
| 162 |
+
return cls(**{key: value for key, value in payload.items() if key in allowed})
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
def haiku_config(**overrides: Any) -> TinyGDNConfig:
|
| 166 |
+
"""650M Haiku: 3×KDA + 1×Gated-MLA, AttnRes, SiTU-GLU, MTP."""
|
| 167 |
+
|
| 168 |
+
payload = {
|
| 169 |
+
"architecture": "HaikuForCausalLM",
|
| 170 |
+
"model_type": "haiku",
|
| 171 |
+
"vocab_size": 65_536,
|
| 172 |
+
"hidden_size": 1024,
|
| 173 |
+
"intermediate_size": 3840,
|
| 174 |
+
"num_hidden_layers": 36,
|
| 175 |
+
"num_attention_heads": 8,
|
| 176 |
+
"num_key_value_heads": 8,
|
| 177 |
+
"attention_head_dim": 128,
|
| 178 |
+
"full_attention_interval": 4,
|
| 179 |
+
"linear_num_heads": 8,
|
| 180 |
+
"linear_num_value_heads": 8,
|
| 181 |
+
"linear_head_dim": 128,
|
| 182 |
+
"linear_layer_type": "kda",
|
| 183 |
+
"attention_layer_type": "gated_mla",
|
| 184 |
+
"mlp_activation": "situ_glu",
|
| 185 |
+
"use_block_attn_res": True,
|
| 186 |
+
"mtp_num_heads": 2,
|
| 187 |
+
"mtp_adapter_rank": 128,
|
| 188 |
+
"mtp_loss_weight": 0.2,
|
| 189 |
+
"max_position_embeddings": 32_768,
|
| 190 |
+
"training_sequence_length": 2048,
|
| 191 |
+
}
|
| 192 |
+
payload.update(overrides)
|
| 193 |
+
return TinyGDNConfig(**payload)
|
tiny_gdn/haiku_layers.py
ADDED
|
@@ -0,0 +1,199 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Haiku mixers: Gated MLA (NoPE), SiTU-GLU, and block Attention Residuals."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import math
|
| 6 |
+
|
| 7 |
+
import torch
|
| 8 |
+
import torch.nn.functional as F
|
| 9 |
+
from torch import nn
|
| 10 |
+
from torch.nn.attention import sdpa_kernel
|
| 11 |
+
|
| 12 |
+
from tiny_gdn.config import TinyGDNConfig
|
| 13 |
+
from tiny_gdn.nn_common import RMSNorm, attention_sdpa_backends
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
class GatedMultiheadLatentAttention(nn.Module):
|
| 17 |
+
"""DeepSeek-style MLA with Kimi K3 NoPE and a full-rank output gate.
|
| 18 |
+
|
| 19 |
+
Queries and keys are content-only. Position is left to the KDA layers.
|
| 20 |
+
KV is compressed through a latent bottleneck, then expanded per head.
|
| 21 |
+
"""
|
| 22 |
+
|
| 23 |
+
def __init__(self, config: TinyGDNConfig) -> None:
|
| 24 |
+
super().__init__()
|
| 25 |
+
self.num_heads = config.num_attention_heads
|
| 26 |
+
self.qk_dim = config.mla_qk_nope_head_dim
|
| 27 |
+
self.v_dim = config.mla_v_head_dim
|
| 28 |
+
self.dropout = config.attention_dropout
|
| 29 |
+
q_out = self.num_heads * self.qk_dim
|
| 30 |
+
v_out = self.num_heads * self.v_dim
|
| 31 |
+
kv_out = self.num_heads * (self.qk_dim + self.v_dim)
|
| 32 |
+
|
| 33 |
+
if config.mla_q_lora_rank is None:
|
| 34 |
+
self.q_proj: nn.Module = nn.Linear(config.hidden_size, q_out, bias=False)
|
| 35 |
+
else:
|
| 36 |
+
self.q_proj = nn.Sequential(
|
| 37 |
+
nn.Linear(config.hidden_size, config.mla_q_lora_rank, bias=False),
|
| 38 |
+
RMSNorm(config.mla_q_lora_rank, config.rms_norm_eps),
|
| 39 |
+
nn.Linear(config.mla_q_lora_rank, q_out, bias=False),
|
| 40 |
+
)
|
| 41 |
+
self.kv_down = nn.Linear(config.hidden_size, config.mla_kv_lora_rank, bias=False)
|
| 42 |
+
self.kv_norm = RMSNorm(config.mla_kv_lora_rank, config.rms_norm_eps)
|
| 43 |
+
self.kv_up = nn.Linear(config.mla_kv_lora_rank, kv_out, bias=False)
|
| 44 |
+
self.q_norm = RMSNorm(self.qk_dim, config.rms_norm_eps)
|
| 45 |
+
self.k_norm = RMSNorm(self.qk_dim, config.rms_norm_eps)
|
| 46 |
+
self.output_gate_proj = nn.Linear(config.hidden_size, v_out, bias=False)
|
| 47 |
+
self.o_proj = nn.Linear(v_out, config.hidden_size, bias=False)
|
| 48 |
+
|
| 49 |
+
def _attention_mask(
|
| 50 |
+
self,
|
| 51 |
+
attention_mask: torch.Tensor | None,
|
| 52 |
+
sequence_length: int,
|
| 53 |
+
device: torch.device,
|
| 54 |
+
) -> torch.Tensor | None:
|
| 55 |
+
if attention_mask is None:
|
| 56 |
+
return None
|
| 57 |
+
if attention_mask.ndim != 2:
|
| 58 |
+
raise ValueError("attention_mask must have shape [batch, sequence]")
|
| 59 |
+
if attention_mask.shape[1] != sequence_length:
|
| 60 |
+
raise ValueError("attention_mask sequence length does not match input")
|
| 61 |
+
causal = torch.ones(
|
| 62 |
+
sequence_length,
|
| 63 |
+
sequence_length,
|
| 64 |
+
dtype=torch.bool,
|
| 65 |
+
device=device,
|
| 66 |
+
).tril()
|
| 67 |
+
valid_keys = attention_mask[:, None, None, :].to(dtype=torch.bool, device=device)
|
| 68 |
+
return causal[None, None, :, :] & valid_keys
|
| 69 |
+
|
| 70 |
+
def _project_kv(self, hidden_states: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
| 71 |
+
batch_size, sequence_length, _ = hidden_states.shape
|
| 72 |
+
compressed = self.kv_norm(self.kv_down(hidden_states))
|
| 73 |
+
key_value = self.kv_up(compressed)
|
| 74 |
+
key, value = key_value.split(
|
| 75 |
+
[self.num_heads * self.qk_dim, self.num_heads * self.v_dim],
|
| 76 |
+
dim=-1,
|
| 77 |
+
)
|
| 78 |
+
key = self.k_norm(
|
| 79 |
+
key.view(batch_size, sequence_length, self.num_heads, self.qk_dim)
|
| 80 |
+
).transpose(1, 2)
|
| 81 |
+
value = value.view(
|
| 82 |
+
batch_size, sequence_length, self.num_heads, self.v_dim
|
| 83 |
+
).transpose(1, 2)
|
| 84 |
+
return key, value
|
| 85 |
+
|
| 86 |
+
def forward(
|
| 87 |
+
self,
|
| 88 |
+
hidden_states: torch.Tensor,
|
| 89 |
+
attention_mask: torch.Tensor | None = None,
|
| 90 |
+
past_key_value: tuple[torch.Tensor, torch.Tensor] | None = None,
|
| 91 |
+
) -> tuple[torch.Tensor, tuple[torch.Tensor, torch.Tensor] | None]:
|
| 92 |
+
batch_size, sequence_length, _ = hidden_states.shape
|
| 93 |
+
query = self.q_proj(hidden_states)
|
| 94 |
+
query = self.q_norm(
|
| 95 |
+
query.view(batch_size, sequence_length, self.num_heads, self.qk_dim)
|
| 96 |
+
).transpose(1, 2)
|
| 97 |
+
key, value = self._project_kv(hidden_states)
|
| 98 |
+
if past_key_value is not None:
|
| 99 |
+
key = torch.cat([past_key_value[0], key], dim=2)
|
| 100 |
+
value = torch.cat([past_key_value[1], value], dim=2)
|
| 101 |
+
present = (key, value)
|
| 102 |
+
past_len = 0 if past_key_value is None else past_key_value[0].shape[2]
|
| 103 |
+
kv_len = key.shape[2]
|
| 104 |
+
|
| 105 |
+
if attention_mask is not None and past_key_value is None:
|
| 106 |
+
sdpa_mask = self._attention_mask(
|
| 107 |
+
attention_mask, sequence_length, hidden_states.device
|
| 108 |
+
)
|
| 109 |
+
is_causal = False
|
| 110 |
+
elif past_key_value is not None and sequence_length == 1:
|
| 111 |
+
sdpa_mask = (
|
| 112 |
+
None
|
| 113 |
+
if attention_mask is None
|
| 114 |
+
else attention_mask[:, None, None, :].to(
|
| 115 |
+
dtype=torch.bool, device=hidden_states.device
|
| 116 |
+
)
|
| 117 |
+
)
|
| 118 |
+
is_causal = False
|
| 119 |
+
elif past_key_value is not None:
|
| 120 |
+
q_idx = torch.arange(
|
| 121 |
+
past_len, past_len + sequence_length, device=hidden_states.device
|
| 122 |
+
)[:, None]
|
| 123 |
+
k_idx = torch.arange(kv_len, device=hidden_states.device)[None, :]
|
| 124 |
+
sdpa_mask = (k_idx <= q_idx)[None, None, :, :]
|
| 125 |
+
is_causal = False
|
| 126 |
+
else:
|
| 127 |
+
sdpa_mask = None
|
| 128 |
+
is_causal = True
|
| 129 |
+
|
| 130 |
+
with sdpa_kernel(attention_sdpa_backends(query.device)):
|
| 131 |
+
attention_output = F.scaled_dot_product_attention(
|
| 132 |
+
query,
|
| 133 |
+
key,
|
| 134 |
+
value,
|
| 135 |
+
attn_mask=sdpa_mask,
|
| 136 |
+
dropout_p=self.dropout if self.training else 0.0,
|
| 137 |
+
is_causal=is_causal,
|
| 138 |
+
)
|
| 139 |
+
attention_output = attention_output.transpose(1, 2).reshape(
|
| 140 |
+
batch_size, sequence_length, -1
|
| 141 |
+
)
|
| 142 |
+
attention_output = attention_output * torch.sigmoid(
|
| 143 |
+
self.output_gate_proj(hidden_states)
|
| 144 |
+
)
|
| 145 |
+
return self.o_proj(attention_output), present
|
| 146 |
+
|
| 147 |
+
|
| 148 |
+
class SiTUGLU(nn.Module):
|
| 149 |
+
"""Sigmoid-Tanh Unit GLU from Kimi K3.
|
| 150 |
+
|
| 151 |
+
Caps the Swish linear factor and the up-projection so the routed / deep
|
| 152 |
+
stack cannot explode, while matching SwiGLU near the origin.
|
| 153 |
+
"""
|
| 154 |
+
|
| 155 |
+
def __init__(self, config: TinyGDNConfig) -> None:
|
| 156 |
+
super().__init__()
|
| 157 |
+
self.gate_up_proj = nn.Linear(
|
| 158 |
+
config.hidden_size,
|
| 159 |
+
config.intermediate_size * 2,
|
| 160 |
+
bias=False,
|
| 161 |
+
)
|
| 162 |
+
self.down_proj = nn.Linear(
|
| 163 |
+
config.intermediate_size,
|
| 164 |
+
config.hidden_size,
|
| 165 |
+
bias=False,
|
| 166 |
+
)
|
| 167 |
+
self.gate_cap = config.situ_gate_cap
|
| 168 |
+
self.up_cap = config.situ_up_cap
|
| 169 |
+
|
| 170 |
+
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
| 171 |
+
gate, up = self.gate_up_proj(hidden_states).chunk(2, dim=-1)
|
| 172 |
+
gated = _softcap(gate, self.gate_cap) * torch.sigmoid(gate)
|
| 173 |
+
return self.down_proj(gated * _softcap(up, self.up_cap))
|
| 174 |
+
|
| 175 |
+
|
| 176 |
+
class BlockAttentionResidual(nn.Module):
|
| 177 |
+
"""Selective depth skip over block outputs (Kimi K3 AttnRes, block-level)."""
|
| 178 |
+
|
| 179 |
+
def __init__(self, hidden_size: int) -> None:
|
| 180 |
+
super().__init__()
|
| 181 |
+
self.query = nn.Parameter(torch.zeros(hidden_size))
|
| 182 |
+
self.scale = 1.0 / math.sqrt(hidden_size)
|
| 183 |
+
|
| 184 |
+
def forward(
|
| 185 |
+
self,
|
| 186 |
+
hidden_states: torch.Tensor,
|
| 187 |
+
memories: list[torch.Tensor],
|
| 188 |
+
) -> torch.Tensor:
|
| 189 |
+
if not memories:
|
| 190 |
+
return hidden_states
|
| 191 |
+
stacked = torch.stack(memories, dim=2)
|
| 192 |
+
scores = torch.einsum("d,bsnd->bsn", self.query, stacked) * self.scale
|
| 193 |
+
weights = torch.softmax(scores, dim=-1)
|
| 194 |
+
mixed = torch.einsum("bsn,bsnd->bsd", weights, stacked)
|
| 195 |
+
return hidden_states + mixed
|
| 196 |
+
|
| 197 |
+
|
| 198 |
+
def _softcap(values: torch.Tensor, cap: float) -> torch.Tensor:
|
| 199 |
+
return cap * torch.tanh(values / cap)
|
tiny_gdn/model.py
ADDED
|
@@ -0,0 +1,713 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import math
|
| 4 |
+
from dataclasses import dataclass
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
from typing import Any
|
| 7 |
+
|
| 8 |
+
import torch
|
| 9 |
+
import torch.nn.functional as F
|
| 10 |
+
from safetensors.torch import load_model, save_model
|
| 11 |
+
from torch import nn
|
| 12 |
+
from torch.nn.attention import sdpa_kernel
|
| 13 |
+
from torch.utils.checkpoint import checkpoint
|
| 14 |
+
|
| 15 |
+
from tiny_gdn.config import TinyGDNConfig
|
| 16 |
+
from tiny_gdn.haiku_layers import (
|
| 17 |
+
BlockAttentionResidual,
|
| 18 |
+
GatedMultiheadLatentAttention,
|
| 19 |
+
SiTUGLU,
|
| 20 |
+
)
|
| 21 |
+
from tiny_gdn.nn_common import RMSNorm, attention_sdpa_backends
|
| 22 |
+
|
| 23 |
+
try:
|
| 24 |
+
# Import the module directly — `from fla.layers import GatedDeltaNet2`
|
| 25 |
+
# executes layers/__init__.py and eagerly loads every attention kernel.
|
| 26 |
+
from fla.layers.gdn2 import GatedDeltaNet2
|
| 27 |
+
except ImportError as import_error:
|
| 28 |
+
GatedDeltaNet2 = None
|
| 29 |
+
FLA_IMPORT_ERROR: ImportError | None = import_error
|
| 30 |
+
else:
|
| 31 |
+
FLA_IMPORT_ERROR = None
|
| 32 |
+
|
| 33 |
+
try:
|
| 34 |
+
from fla.layers.kda import KimiDeltaAttention
|
| 35 |
+
except ImportError as import_error:
|
| 36 |
+
KimiDeltaAttention = None
|
| 37 |
+
KDA_IMPORT_ERROR: ImportError | None = import_error
|
| 38 |
+
else:
|
| 39 |
+
KDA_IMPORT_ERROR = None
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
@dataclass
|
| 43 |
+
class TinyGDNOutput:
|
| 44 |
+
loss: torch.Tensor | None
|
| 45 |
+
logits: torch.Tensor | None
|
| 46 |
+
main_loss: torch.Tensor | None
|
| 47 |
+
mtp_loss: torch.Tensor | None
|
| 48 |
+
z_loss: torch.Tensor | None
|
| 49 |
+
hidden_states: torch.Tensor | None = None
|
| 50 |
+
past_key_values: Any | None = None
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
class RotaryEmbedding(nn.Module):
|
| 54 |
+
def __init__(self, rotary_dim: int, rope_theta: float) -> None:
|
| 55 |
+
super().__init__()
|
| 56 |
+
inverse_frequency = 1.0 / (
|
| 57 |
+
rope_theta
|
| 58 |
+
** (
|
| 59 |
+
torch.arange(0, rotary_dim, 2, dtype=torch.float32)
|
| 60 |
+
/ rotary_dim
|
| 61 |
+
)
|
| 62 |
+
)
|
| 63 |
+
self.rotary_dim = rotary_dim
|
| 64 |
+
self.register_buffer("inverse_frequency", inverse_frequency, persistent=False)
|
| 65 |
+
|
| 66 |
+
def forward(
|
| 67 |
+
self,
|
| 68 |
+
sequence_length: int,
|
| 69 |
+
device: torch.device,
|
| 70 |
+
dtype: torch.dtype,
|
| 71 |
+
position_offset: int = 0,
|
| 72 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 73 |
+
positions = torch.arange(
|
| 74 |
+
position_offset,
|
| 75 |
+
position_offset + sequence_length,
|
| 76 |
+
device=device,
|
| 77 |
+
dtype=torch.float32,
|
| 78 |
+
)
|
| 79 |
+
frequencies = torch.outer(positions, self.inverse_frequency.float())
|
| 80 |
+
embeddings = torch.cat((frequencies, frequencies), dim=-1)
|
| 81 |
+
return embeddings.cos().to(dtype=dtype), embeddings.sin().to(dtype=dtype)
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
def rotate_half(hidden_states: torch.Tensor) -> torch.Tensor:
|
| 85 |
+
first, second = hidden_states.chunk(2, dim=-1)
|
| 86 |
+
return torch.cat((-second, first), dim=-1)
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
def apply_rotary_embedding(
|
| 90 |
+
query: torch.Tensor,
|
| 91 |
+
key: torch.Tensor,
|
| 92 |
+
cosine: torch.Tensor,
|
| 93 |
+
sine: torch.Tensor,
|
| 94 |
+
rotary_dim: int,
|
| 95 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 96 |
+
cosine = cosine[None, None, :, :]
|
| 97 |
+
sine = sine[None, None, :, :]
|
| 98 |
+
query_rotary, query_pass = query[..., :rotary_dim], query[..., rotary_dim:]
|
| 99 |
+
key_rotary, key_pass = key[..., :rotary_dim], key[..., rotary_dim:]
|
| 100 |
+
query_rotary = query_rotary * cosine + rotate_half(query_rotary) * sine
|
| 101 |
+
key_rotary = key_rotary * cosine + rotate_half(key_rotary) * sine
|
| 102 |
+
return (
|
| 103 |
+
torch.cat((query_rotary, query_pass), dim=-1),
|
| 104 |
+
torch.cat((key_rotary, key_pass), dim=-1),
|
| 105 |
+
)
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
class GatedGroupedQueryAttention(nn.Module):
|
| 109 |
+
"""QK-normalized, partially rotary GQA with a learned sigmoid output gate."""
|
| 110 |
+
|
| 111 |
+
def __init__(self, config: TinyGDNConfig) -> None:
|
| 112 |
+
super().__init__()
|
| 113 |
+
self.num_heads = config.num_attention_heads
|
| 114 |
+
self.num_key_value_heads = config.num_key_value_heads
|
| 115 |
+
self.head_dim = config.attention_head_dim
|
| 116 |
+
self.rotary_dim = config.rotary_dim
|
| 117 |
+
self.dropout = config.attention_dropout
|
| 118 |
+
|
| 119 |
+
query_size = self.num_heads * self.head_dim
|
| 120 |
+
key_value_size = self.num_key_value_heads * self.head_dim
|
| 121 |
+
self.q_gate_proj = nn.Linear(config.hidden_size, query_size * 2, bias=False)
|
| 122 |
+
self.k_proj = nn.Linear(config.hidden_size, key_value_size, bias=False)
|
| 123 |
+
self.v_proj = nn.Linear(config.hidden_size, key_value_size, bias=False)
|
| 124 |
+
self.o_proj = nn.Linear(query_size, config.hidden_size, bias=False)
|
| 125 |
+
self.q_norm = RMSNorm(self.head_dim, config.rms_norm_eps)
|
| 126 |
+
self.k_norm = RMSNorm(self.head_dim, config.rms_norm_eps)
|
| 127 |
+
self.rotary = RotaryEmbedding(self.rotary_dim, config.rope_theta)
|
| 128 |
+
|
| 129 |
+
def _attention_mask(
|
| 130 |
+
self,
|
| 131 |
+
attention_mask: torch.Tensor | None,
|
| 132 |
+
sequence_length: int,
|
| 133 |
+
device: torch.device,
|
| 134 |
+
) -> torch.Tensor | None:
|
| 135 |
+
if attention_mask is None:
|
| 136 |
+
return None
|
| 137 |
+
if attention_mask.ndim != 2:
|
| 138 |
+
raise ValueError("attention_mask must have shape [batch, sequence]")
|
| 139 |
+
if attention_mask.shape[1] != sequence_length:
|
| 140 |
+
raise ValueError("attention_mask sequence length does not match input")
|
| 141 |
+
|
| 142 |
+
causal = torch.ones(
|
| 143 |
+
sequence_length,
|
| 144 |
+
sequence_length,
|
| 145 |
+
dtype=torch.bool,
|
| 146 |
+
device=device,
|
| 147 |
+
).tril()
|
| 148 |
+
valid_keys = attention_mask[:, None, None, :].to(dtype=torch.bool, device=device)
|
| 149 |
+
return causal[None, None, :, :] & valid_keys
|
| 150 |
+
|
| 151 |
+
def forward(
|
| 152 |
+
self,
|
| 153 |
+
hidden_states: torch.Tensor,
|
| 154 |
+
attention_mask: torch.Tensor | None = None,
|
| 155 |
+
past_key_value: tuple[torch.Tensor, torch.Tensor] | None = None,
|
| 156 |
+
) -> tuple[torch.Tensor, tuple[torch.Tensor, torch.Tensor] | None]:
|
| 157 |
+
batch_size, sequence_length, _ = hidden_states.shape
|
| 158 |
+
past_len = 0 if past_key_value is None else past_key_value[0].shape[2]
|
| 159 |
+
query_and_gate = self.q_gate_proj(hidden_states)
|
| 160 |
+
query, output_gate = query_and_gate.chunk(2, dim=-1)
|
| 161 |
+
|
| 162 |
+
query = query.view(batch_size, sequence_length, self.num_heads, self.head_dim)
|
| 163 |
+
key = self.k_proj(hidden_states).view(
|
| 164 |
+
batch_size,
|
| 165 |
+
sequence_length,
|
| 166 |
+
self.num_key_value_heads,
|
| 167 |
+
self.head_dim,
|
| 168 |
+
)
|
| 169 |
+
value = self.v_proj(hidden_states).view(
|
| 170 |
+
batch_size,
|
| 171 |
+
sequence_length,
|
| 172 |
+
self.num_key_value_heads,
|
| 173 |
+
self.head_dim,
|
| 174 |
+
)
|
| 175 |
+
|
| 176 |
+
query = self.q_norm(query).transpose(1, 2)
|
| 177 |
+
key = self.k_norm(key).transpose(1, 2)
|
| 178 |
+
value = value.transpose(1, 2)
|
| 179 |
+
|
| 180 |
+
cosine, sine = self.rotary(
|
| 181 |
+
sequence_length,
|
| 182 |
+
device=hidden_states.device,
|
| 183 |
+
dtype=query.dtype,
|
| 184 |
+
position_offset=past_len,
|
| 185 |
+
)
|
| 186 |
+
query, key = apply_rotary_embedding(
|
| 187 |
+
query,
|
| 188 |
+
key,
|
| 189 |
+
cosine,
|
| 190 |
+
sine,
|
| 191 |
+
rotary_dim=self.rotary_dim,
|
| 192 |
+
)
|
| 193 |
+
if past_key_value is not None:
|
| 194 |
+
key = torch.cat([past_key_value[0], key], dim=2)
|
| 195 |
+
value = torch.cat([past_key_value[1], value], dim=2)
|
| 196 |
+
present = (key, value)
|
| 197 |
+
|
| 198 |
+
kv_len = key.shape[2]
|
| 199 |
+
if attention_mask is not None and past_key_value is None:
|
| 200 |
+
sdpa_mask = self._attention_mask(
|
| 201 |
+
attention_mask,
|
| 202 |
+
sequence_length,
|
| 203 |
+
hidden_states.device,
|
| 204 |
+
)
|
| 205 |
+
is_causal = False
|
| 206 |
+
elif past_key_value is not None and sequence_length == 1:
|
| 207 |
+
# Decode step: query attends to the cached key/value prefix. Preserve
|
| 208 |
+
# the prefill padding mask when decoding a left-padded prompt batch.
|
| 209 |
+
sdpa_mask = (
|
| 210 |
+
None
|
| 211 |
+
if attention_mask is None
|
| 212 |
+
else attention_mask[:, None, None, :].to(
|
| 213 |
+
dtype=torch.bool,
|
| 214 |
+
device=hidden_states.device,
|
| 215 |
+
)
|
| 216 |
+
)
|
| 217 |
+
is_causal = False
|
| 218 |
+
elif past_key_value is not None:
|
| 219 |
+
# Prefill chunk with cache — build causal mask over kv_len.
|
| 220 |
+
q_idx = torch.arange(
|
| 221 |
+
past_len, past_len + sequence_length, device=hidden_states.device
|
| 222 |
+
)[:, None]
|
| 223 |
+
k_idx = torch.arange(kv_len, device=hidden_states.device)[None, :]
|
| 224 |
+
sdpa_mask = (k_idx <= q_idx)[None, None, :, :]
|
| 225 |
+
is_causal = False
|
| 226 |
+
else:
|
| 227 |
+
sdpa_mask = None
|
| 228 |
+
is_causal = True
|
| 229 |
+
sdpa_options = {
|
| 230 |
+
"attn_mask": sdpa_mask,
|
| 231 |
+
"dropout_p": self.dropout if self.training else 0.0,
|
| 232 |
+
"is_causal": is_causal,
|
| 233 |
+
"enable_gqa": True,
|
| 234 |
+
}
|
| 235 |
+
with sdpa_kernel(attention_sdpa_backends(query.device)):
|
| 236 |
+
attention_output = F.scaled_dot_product_attention(
|
| 237 |
+
query,
|
| 238 |
+
key,
|
| 239 |
+
value,
|
| 240 |
+
**sdpa_options,
|
| 241 |
+
)
|
| 242 |
+
attention_output = attention_output.transpose(1, 2).reshape(
|
| 243 |
+
batch_size,
|
| 244 |
+
sequence_length,
|
| 245 |
+
-1,
|
| 246 |
+
)
|
| 247 |
+
attention_output = attention_output * torch.sigmoid(output_gate)
|
| 248 |
+
return self.o_proj(attention_output), present
|
| 249 |
+
|
| 250 |
+
|
| 251 |
+
class SwiGLU(nn.Module):
|
| 252 |
+
def __init__(self, config: TinyGDNConfig) -> None:
|
| 253 |
+
super().__init__()
|
| 254 |
+
self.gate_up_proj = nn.Linear(
|
| 255 |
+
config.hidden_size,
|
| 256 |
+
config.intermediate_size * 2,
|
| 257 |
+
bias=False,
|
| 258 |
+
)
|
| 259 |
+
self.down_proj = nn.Linear(
|
| 260 |
+
config.intermediate_size,
|
| 261 |
+
config.hidden_size,
|
| 262 |
+
bias=False,
|
| 263 |
+
)
|
| 264 |
+
|
| 265 |
+
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
| 266 |
+
gate, up = self.gate_up_proj(hidden_states).chunk(2, dim=-1)
|
| 267 |
+
return self.down_proj(F.silu(gate) * up)
|
| 268 |
+
|
| 269 |
+
|
| 270 |
+
class TinyGDNBlock(nn.Module):
|
| 271 |
+
def __init__(self, config: TinyGDNConfig, layer_index: int) -> None:
|
| 272 |
+
super().__init__()
|
| 273 |
+
layer_type = config.layer_types[layer_index]
|
| 274 |
+
self.layer_type = layer_type
|
| 275 |
+
self.token_mixer_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
| 276 |
+
self.mlp_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
| 277 |
+
|
| 278 |
+
if layer_type == "gdn2":
|
| 279 |
+
if GatedDeltaNet2 is None:
|
| 280 |
+
raise ImportError(
|
| 281 |
+
"Gated DeltaNet-2 requires the pinned flash-linear-attention dependency"
|
| 282 |
+
) from FLA_IMPORT_ERROR
|
| 283 |
+
self.token_mixer = GatedDeltaNet2(
|
| 284 |
+
hidden_size=config.hidden_size,
|
| 285 |
+
expand_v=config.linear_expand_v,
|
| 286 |
+
head_dim=config.linear_head_dim,
|
| 287 |
+
num_heads=config.linear_num_heads,
|
| 288 |
+
num_v_heads=config.linear_num_value_heads,
|
| 289 |
+
mode="chunk",
|
| 290 |
+
use_short_conv=True,
|
| 291 |
+
allow_neg_eigval=config.allow_negative_eigenvalues,
|
| 292 |
+
conv_size=config.linear_conv_kernel_dim,
|
| 293 |
+
conv_bias=False,
|
| 294 |
+
layer_idx=layer_index,
|
| 295 |
+
norm_eps=config.rms_norm_eps,
|
| 296 |
+
)
|
| 297 |
+
elif layer_type == "kda":
|
| 298 |
+
if KimiDeltaAttention is None:
|
| 299 |
+
raise ImportError(
|
| 300 |
+
"Kimi Delta Attention requires the pinned flash-linear-attention dependency"
|
| 301 |
+
) from KDA_IMPORT_ERROR
|
| 302 |
+
self.token_mixer = KimiDeltaAttention(
|
| 303 |
+
hidden_size=config.hidden_size,
|
| 304 |
+
expand_v=config.linear_expand_v,
|
| 305 |
+
head_dim=config.linear_head_dim,
|
| 306 |
+
num_heads=config.linear_num_heads,
|
| 307 |
+
num_v_heads=config.linear_num_value_heads,
|
| 308 |
+
mode="chunk",
|
| 309 |
+
use_short_conv=True,
|
| 310 |
+
allow_neg_eigval=config.allow_negative_eigenvalues,
|
| 311 |
+
safe_gate=config.kda_safe_gate,
|
| 312 |
+
lower_bound=config.kda_lower_bound,
|
| 313 |
+
conv_size=config.linear_conv_kernel_dim,
|
| 314 |
+
conv_bias=False,
|
| 315 |
+
layer_idx=layer_index,
|
| 316 |
+
norm_eps=config.rms_norm_eps,
|
| 317 |
+
)
|
| 318 |
+
elif layer_type == "full_attention":
|
| 319 |
+
self.token_mixer = GatedGroupedQueryAttention(config)
|
| 320 |
+
elif layer_type == "gated_mla":
|
| 321 |
+
self.token_mixer = GatedMultiheadLatentAttention(config)
|
| 322 |
+
else:
|
| 323 |
+
raise ValueError(f"Unsupported layer type: {layer_type}")
|
| 324 |
+
|
| 325 |
+
if config.mlp_activation == "swiglu":
|
| 326 |
+
self.mlp: nn.Module = SwiGLU(config)
|
| 327 |
+
elif config.mlp_activation == "situ_glu":
|
| 328 |
+
self.mlp = SiTUGLU(config)
|
| 329 |
+
else:
|
| 330 |
+
raise ValueError(f"Unsupported mlp_activation: {config.mlp_activation}")
|
| 331 |
+
|
| 332 |
+
def forward(
|
| 333 |
+
self,
|
| 334 |
+
hidden_states: torch.Tensor,
|
| 335 |
+
attention_mask: torch.Tensor | None = None,
|
| 336 |
+
*,
|
| 337 |
+
past_key_values: Any | None = None,
|
| 338 |
+
past_key_value: tuple[torch.Tensor, torch.Tensor] | None = None,
|
| 339 |
+
use_cache: bool = False,
|
| 340 |
+
) -> tuple[torch.Tensor, Any]:
|
| 341 |
+
residual = hidden_states
|
| 342 |
+
normalized = self.token_mixer_norm(hidden_states)
|
| 343 |
+
present: Any = None
|
| 344 |
+
if self.layer_type in {"gdn2", "kda"}:
|
| 345 |
+
mixed, _, past_key_values = self.token_mixer(
|
| 346 |
+
normalized,
|
| 347 |
+
attention_mask=attention_mask,
|
| 348 |
+
past_key_values=past_key_values,
|
| 349 |
+
use_cache=use_cache,
|
| 350 |
+
)
|
| 351 |
+
present = past_key_values
|
| 352 |
+
elif self.layer_type in {"full_attention", "gated_mla"}:
|
| 353 |
+
mixed, present = self.token_mixer(
|
| 354 |
+
normalized,
|
| 355 |
+
attention_mask=attention_mask,
|
| 356 |
+
past_key_value=past_key_value,
|
| 357 |
+
)
|
| 358 |
+
if not use_cache:
|
| 359 |
+
present = None
|
| 360 |
+
else:
|
| 361 |
+
raise ValueError(f"Unsupported layer type: {self.layer_type}")
|
| 362 |
+
hidden_states = residual + mixed
|
| 363 |
+
hidden_states = hidden_states + self.mlp(self.mlp_norm(hidden_states))
|
| 364 |
+
return hidden_states, present
|
| 365 |
+
|
| 366 |
+
|
| 367 |
+
class MultiTokenPredictionAdapter(nn.Module):
|
| 368 |
+
"""A lightweight residual adapter for one additional prediction horizon."""
|
| 369 |
+
|
| 370 |
+
def __init__(self, config: TinyGDNConfig) -> None:
|
| 371 |
+
super().__init__()
|
| 372 |
+
self.norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
| 373 |
+
self.down_proj = nn.Linear(
|
| 374 |
+
config.hidden_size,
|
| 375 |
+
config.mtp_adapter_rank,
|
| 376 |
+
bias=False,
|
| 377 |
+
)
|
| 378 |
+
self.up_proj = nn.Linear(
|
| 379 |
+
config.mtp_adapter_rank,
|
| 380 |
+
config.hidden_size,
|
| 381 |
+
bias=False,
|
| 382 |
+
)
|
| 383 |
+
|
| 384 |
+
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
| 385 |
+
adapted = self.up_proj(F.silu(self.down_proj(self.norm(hidden_states))))
|
| 386 |
+
return hidden_states + adapted
|
| 387 |
+
|
| 388 |
+
|
| 389 |
+
class TinyGDNForCausalLM(nn.Module):
|
| 390 |
+
def __init__(self, config: TinyGDNConfig) -> None:
|
| 391 |
+
super().__init__()
|
| 392 |
+
self.config = config
|
| 393 |
+
self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size)
|
| 394 |
+
self.layers = nn.ModuleList(
|
| 395 |
+
TinyGDNBlock(config, layer_index)
|
| 396 |
+
for layer_index in range(config.num_hidden_layers)
|
| 397 |
+
)
|
| 398 |
+
self.final_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
| 399 |
+
self.mtp_adapters = nn.ModuleList(
|
| 400 |
+
MultiTokenPredictionAdapter(config)
|
| 401 |
+
for _ in range(config.mtp_num_heads)
|
| 402 |
+
)
|
| 403 |
+
self.block_attn_res = (
|
| 404 |
+
BlockAttentionResidual(config.hidden_size)
|
| 405 |
+
if config.use_block_attn_res
|
| 406 |
+
else None
|
| 407 |
+
)
|
| 408 |
+
self.gradient_checkpointing = False
|
| 409 |
+
|
| 410 |
+
self.apply(self._initialize_module)
|
| 411 |
+
self._initialize_residual_projections()
|
| 412 |
+
|
| 413 |
+
def _initialize_module(self, module: nn.Module) -> None:
|
| 414 |
+
if isinstance(module, nn.Linear):
|
| 415 |
+
nn.init.normal_(
|
| 416 |
+
module.weight,
|
| 417 |
+
mean=0.0,
|
| 418 |
+
std=self.config.initializer_range,
|
| 419 |
+
)
|
| 420 |
+
if module.bias is not None:
|
| 421 |
+
nn.init.zeros_(module.bias)
|
| 422 |
+
elif isinstance(module, nn.Embedding):
|
| 423 |
+
nn.init.normal_(
|
| 424 |
+
module.weight,
|
| 425 |
+
mean=0.0,
|
| 426 |
+
std=self.config.initializer_range,
|
| 427 |
+
)
|
| 428 |
+
|
| 429 |
+
def _initialize_residual_projections(self) -> None:
|
| 430 |
+
residual_std = self.config.initializer_range / math.sqrt(
|
| 431 |
+
2 * self.config.num_hidden_layers
|
| 432 |
+
)
|
| 433 |
+
for layer in self.layers:
|
| 434 |
+
nn.init.normal_(
|
| 435 |
+
layer.token_mixer.o_proj.weight,
|
| 436 |
+
mean=0.0,
|
| 437 |
+
std=residual_std,
|
| 438 |
+
)
|
| 439 |
+
nn.init.normal_(
|
| 440 |
+
layer.mlp.down_proj.weight,
|
| 441 |
+
mean=0.0,
|
| 442 |
+
std=residual_std,
|
| 443 |
+
)
|
| 444 |
+
for adapter in self.mtp_adapters:
|
| 445 |
+
nn.init.normal_(adapter.up_proj.weight, mean=0.0, std=residual_std)
|
| 446 |
+
|
| 447 |
+
def enable_gradient_checkpointing(self, enabled: bool = True) -> None:
|
| 448 |
+
self.gradient_checkpointing = enabled
|
| 449 |
+
|
| 450 |
+
def project_to_vocabulary(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
| 451 |
+
return F.linear(hidden_states, self.embed_tokens.weight)
|
| 452 |
+
|
| 453 |
+
def _run_layer(
|
| 454 |
+
self,
|
| 455 |
+
layer: TinyGDNBlock,
|
| 456 |
+
hidden_states: torch.Tensor,
|
| 457 |
+
attention_mask: torch.Tensor | None,
|
| 458 |
+
*,
|
| 459 |
+
past_key_values: Any | None = None,
|
| 460 |
+
past_key_value: tuple[torch.Tensor, torch.Tensor] | None = None,
|
| 461 |
+
use_cache: bool = False,
|
| 462 |
+
) -> tuple[torch.Tensor, Any]:
|
| 463 |
+
if self.gradient_checkpointing and self.training:
|
| 464 |
+
hidden_states, present = checkpoint(
|
| 465 |
+
layer,
|
| 466 |
+
hidden_states,
|
| 467 |
+
attention_mask,
|
| 468 |
+
use_reentrant=False,
|
| 469 |
+
)
|
| 470 |
+
return hidden_states, present
|
| 471 |
+
return layer(
|
| 472 |
+
hidden_states,
|
| 473 |
+
attention_mask,
|
| 474 |
+
past_key_values=past_key_values,
|
| 475 |
+
past_key_value=past_key_value,
|
| 476 |
+
use_cache=use_cache,
|
| 477 |
+
)
|
| 478 |
+
|
| 479 |
+
def _causal_loss(
|
| 480 |
+
self,
|
| 481 |
+
hidden_states: torch.Tensor,
|
| 482 |
+
labels: torch.Tensor,
|
| 483 |
+
target_offset: int,
|
| 484 |
+
adapter: nn.Module | None = None,
|
| 485 |
+
compute_z_loss: bool = False,
|
| 486 |
+
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
| 487 |
+
if target_offset < 0:
|
| 488 |
+
raise ValueError("target_offset cannot be negative")
|
| 489 |
+
if target_offset and hidden_states.shape[1] <= target_offset:
|
| 490 |
+
raise ValueError(
|
| 491 |
+
f"Sequence length must exceed target offset {target_offset}"
|
| 492 |
+
)
|
| 493 |
+
if target_offset:
|
| 494 |
+
prediction_states = hidden_states[:, :-target_offset, :]
|
| 495 |
+
targets = labels[:, target_offset:].contiguous()
|
| 496 |
+
else:
|
| 497 |
+
prediction_states = hidden_states
|
| 498 |
+
targets = labels.contiguous()
|
| 499 |
+
if adapter is not None:
|
| 500 |
+
prediction_states = adapter(prediction_states)
|
| 501 |
+
logits = self.project_to_vocabulary(prediction_states)
|
| 502 |
+
cross_entropy = F.cross_entropy(
|
| 503 |
+
logits.reshape(-1, self.config.vocab_size),
|
| 504 |
+
targets.reshape(-1),
|
| 505 |
+
ignore_index=-100,
|
| 506 |
+
)
|
| 507 |
+
z_loss = None
|
| 508 |
+
if compute_z_loss:
|
| 509 |
+
valid_targets = targets.ne(-100)
|
| 510 |
+
# Keep logits in their training dtype for logsumexp. Casting the
|
| 511 |
+
# full [B, S, V] tensor to fp32 materializes ~3 GiB in the autograd
|
| 512 |
+
# graph at 16k; upcast only the reduced [B, S] partition.
|
| 513 |
+
log_partition = torch.logsumexp(logits, dim=-1).float()
|
| 514 |
+
z_loss = log_partition.square()[valid_targets].mean()
|
| 515 |
+
return cross_entropy, z_loss
|
| 516 |
+
|
| 517 |
+
def forward(
|
| 518 |
+
self,
|
| 519 |
+
input_ids: torch.Tensor,
|
| 520 |
+
labels: torch.Tensor | None = None,
|
| 521 |
+
attention_mask: torch.Tensor | None = None,
|
| 522 |
+
past_key_values: Any | None = None,
|
| 523 |
+
*,
|
| 524 |
+
use_cache: bool = False,
|
| 525 |
+
return_logits: bool = True,
|
| 526 |
+
return_hidden_states: bool = False,
|
| 527 |
+
labels_are_shifted: bool = False,
|
| 528 |
+
include_mtp_loss: bool = True,
|
| 529 |
+
mtp_loss_weight: float | None = None,
|
| 530 |
+
z_loss_coefficient: float = 0.0,
|
| 531 |
+
logits_to_keep: int | None = None,
|
| 532 |
+
) -> TinyGDNOutput:
|
| 533 |
+
if input_ids.ndim != 2:
|
| 534 |
+
raise ValueError("input_ids must have shape [batch, sequence]")
|
| 535 |
+
if input_ids.shape[1] > self.config.max_position_embeddings:
|
| 536 |
+
raise ValueError("Input exceeds max_position_embeddings")
|
| 537 |
+
if labels is not None and labels.shape != input_ids.shape:
|
| 538 |
+
raise ValueError("labels must have the same shape as input_ids")
|
| 539 |
+
if z_loss_coefficient < 0.0:
|
| 540 |
+
raise ValueError("z_loss_coefficient cannot be negative")
|
| 541 |
+
if logits_to_keep is not None and logits_to_keep <= 0:
|
| 542 |
+
raise ValueError("logits_to_keep must be positive")
|
| 543 |
+
if use_cache and labels is not None:
|
| 544 |
+
raise ValueError("use_cache is not supported with labels")
|
| 545 |
+
effective_mtp_weight = (
|
| 546 |
+
self.config.mtp_loss_weight
|
| 547 |
+
if mtp_loss_weight is None
|
| 548 |
+
else mtp_loss_weight
|
| 549 |
+
)
|
| 550 |
+
if not 0.0 <= effective_mtp_weight <= 1.0:
|
| 551 |
+
raise ValueError("mtp_loss_weight must be between zero and one")
|
| 552 |
+
|
| 553 |
+
if use_cache and past_key_values is None:
|
| 554 |
+
try:
|
| 555 |
+
from fla.models.utils import Cache as FlaCache
|
| 556 |
+
except ImportError as import_error:
|
| 557 |
+
raise ImportError(
|
| 558 |
+
"Cached decode requires flash-linear-attention Cache"
|
| 559 |
+
) from import_error
|
| 560 |
+
past_key_values = {
|
| 561 |
+
"fla": FlaCache(),
|
| 562 |
+
"gqa": [None] * len(self.layers),
|
| 563 |
+
}
|
| 564 |
+
elif past_key_values is not None and not isinstance(past_key_values, dict):
|
| 565 |
+
raise TypeError("past_key_values must be a TinyGDN cache dict or None")
|
| 566 |
+
|
| 567 |
+
fla_cache = None if past_key_values is None else past_key_values["fla"]
|
| 568 |
+
gqa_cache = None if past_key_values is None else past_key_values["gqa"]
|
| 569 |
+
|
| 570 |
+
hidden_states = self.embed_tokens(input_ids)
|
| 571 |
+
shared_layer_indices = set(self.config.shared_layer_indices)
|
| 572 |
+
softmax_types = {"full_attention", "gated_mla"}
|
| 573 |
+
attn_memories: list[torch.Tensor] = [hidden_states]
|
| 574 |
+
interval = self.config.full_attention_interval
|
| 575 |
+
for layer_index, layer in enumerate(self.layers):
|
| 576 |
+
layer_gqa = None if gqa_cache is None else gqa_cache[layer_index]
|
| 577 |
+
hidden_states, present = self._run_layer(
|
| 578 |
+
layer,
|
| 579 |
+
hidden_states,
|
| 580 |
+
attention_mask,
|
| 581 |
+
past_key_values=fla_cache,
|
| 582 |
+
past_key_value=layer_gqa,
|
| 583 |
+
use_cache=use_cache,
|
| 584 |
+
)
|
| 585 |
+
if use_cache and layer.layer_type in softmax_types and gqa_cache is not None:
|
| 586 |
+
gqa_cache[layer_index] = present
|
| 587 |
+
if layer_index in shared_layer_indices:
|
| 588 |
+
layer_gqa = None if gqa_cache is None else gqa_cache[layer_index]
|
| 589 |
+
hidden_states, present = self._run_layer(
|
| 590 |
+
layer,
|
| 591 |
+
hidden_states,
|
| 592 |
+
attention_mask,
|
| 593 |
+
past_key_values=fla_cache,
|
| 594 |
+
past_key_value=layer_gqa,
|
| 595 |
+
use_cache=use_cache,
|
| 596 |
+
)
|
| 597 |
+
if use_cache and layer.layer_type in softmax_types and gqa_cache is not None:
|
| 598 |
+
gqa_cache[layer_index] = present
|
| 599 |
+
if (
|
| 600 |
+
self.block_attn_res is not None
|
| 601 |
+
and (layer_index + 1) % interval == 0
|
| 602 |
+
):
|
| 603 |
+
hidden_states = self.block_attn_res(
|
| 604 |
+
hidden_states, [*attn_memories, hidden_states]
|
| 605 |
+
)
|
| 606 |
+
attn_memories.append(hidden_states)
|
| 607 |
+
hidden_states = self.final_norm(hidden_states)
|
| 608 |
+
|
| 609 |
+
main_loss = None
|
| 610 |
+
mtp_loss = None
|
| 611 |
+
z_loss = None
|
| 612 |
+
total_loss = None
|
| 613 |
+
if labels is not None:
|
| 614 |
+
main_target_offset = 0 if labels_are_shifted else 1
|
| 615 |
+
main_loss, z_loss = self._causal_loss(
|
| 616 |
+
hidden_states,
|
| 617 |
+
labels,
|
| 618 |
+
target_offset=main_target_offset,
|
| 619 |
+
compute_z_loss=z_loss_coefficient > 0.0,
|
| 620 |
+
)
|
| 621 |
+
if self.mtp_adapters and include_mtp_loss:
|
| 622 |
+
auxiliary_losses = [
|
| 623 |
+
self._causal_loss(
|
| 624 |
+
hidden_states,
|
| 625 |
+
labels,
|
| 626 |
+
target_offset=(
|
| 627 |
+
head_index + 1
|
| 628 |
+
if labels_are_shifted
|
| 629 |
+
else head_index + 2
|
| 630 |
+
),
|
| 631 |
+
adapter=adapter,
|
| 632 |
+
)[0]
|
| 633 |
+
for head_index, adapter in enumerate(self.mtp_adapters)
|
| 634 |
+
]
|
| 635 |
+
mtp_loss = torch.stack(auxiliary_losses).mean()
|
| 636 |
+
total_loss = main_loss + effective_mtp_weight * mtp_loss
|
| 637 |
+
else:
|
| 638 |
+
total_loss = main_loss
|
| 639 |
+
if z_loss is not None:
|
| 640 |
+
total_loss = total_loss + z_loss_coefficient * z_loss
|
| 641 |
+
|
| 642 |
+
output_states = (
|
| 643 |
+
hidden_states
|
| 644 |
+
if logits_to_keep is None
|
| 645 |
+
else hidden_states[:, -logits_to_keep:, :]
|
| 646 |
+
)
|
| 647 |
+
logits = self.project_to_vocabulary(output_states) if return_logits else None
|
| 648 |
+
return TinyGDNOutput(
|
| 649 |
+
loss=total_loss,
|
| 650 |
+
logits=logits,
|
| 651 |
+
main_loss=main_loss,
|
| 652 |
+
mtp_loss=mtp_loss,
|
| 653 |
+
z_loss=z_loss,
|
| 654 |
+
hidden_states=hidden_states if return_hidden_states else None,
|
| 655 |
+
past_key_values=past_key_values if use_cache else None,
|
| 656 |
+
)
|
| 657 |
+
|
| 658 |
+
def parameter_report(self) -> dict[str, int]:
|
| 659 |
+
total = sum(parameter.numel() for parameter in self.parameters())
|
| 660 |
+
mtp = sum(parameter.numel() for parameter in self.mtp_adapters.parameters())
|
| 661 |
+
embeddings = self.embed_tokens.weight.numel()
|
| 662 |
+
return {
|
| 663 |
+
"deployable_core": total - mtp,
|
| 664 |
+
"training_total": total,
|
| 665 |
+
"embedding": embeddings,
|
| 666 |
+
"mtp_auxiliary": mtp,
|
| 667 |
+
"non_embedding_core": total - mtp - embeddings,
|
| 668 |
+
}
|
| 669 |
+
|
| 670 |
+
def save_checkpoint(self, output_dir: Path) -> None:
|
| 671 |
+
output_dir.mkdir(parents=True, exist_ok=True)
|
| 672 |
+
self.config.save_json(output_dir / "config.json")
|
| 673 |
+
save_model(self, output_dir / "model.safetensors")
|
| 674 |
+
|
| 675 |
+
@classmethod
|
| 676 |
+
def from_checkpoint(
|
| 677 |
+
cls,
|
| 678 |
+
checkpoint_dir: Path,
|
| 679 |
+
*,
|
| 680 |
+
device: str | torch.device = "cpu",
|
| 681 |
+
dtype: torch.dtype | None = None,
|
| 682 |
+
) -> TinyGDNForCausalLM:
|
| 683 |
+
config = TinyGDNConfig.from_json(checkpoint_dir / "config.json")
|
| 684 |
+
model = cls(config).to(device=device, dtype=dtype)
|
| 685 |
+
load_model(model, checkpoint_dir / "model.safetensors", device=str(device))
|
| 686 |
+
return model
|
| 687 |
+
|
| 688 |
+
def extra_repr(self) -> str:
|
| 689 |
+
report = self.parameter_report()
|
| 690 |
+
return (
|
| 691 |
+
f"core_parameters={report['deployable_core']:,}, "
|
| 692 |
+
f"training_parameters={report['training_total']:,}"
|
| 693 |
+
)
|
| 694 |
+
|
| 695 |
+
def get_architecture_metadata(self) -> dict[str, Any]:
|
| 696 |
+
return {
|
| 697 |
+
"architecture": self.config.architecture,
|
| 698 |
+
"layer_types": list(self.config.layer_types),
|
| 699 |
+
"effective_num_layers": self.config.effective_num_layers,
|
| 700 |
+
"shared_layer_indices": list(self.config.shared_layer_indices),
|
| 701 |
+
"parameter_report": self.parameter_report(),
|
| 702 |
+
"features": [
|
| 703 |
+
"hybrid linear + softmax token mixing",
|
| 704 |
+
"Kimi Delta Attention or Gated DeltaNet-2",
|
| 705 |
+
"gated MLA or gated GQA",
|
| 706 |
+
"optional block Attention Residuals",
|
| 707 |
+
"SwiGLU or SiTU-GLU",
|
| 708 |
+
"QK normalization",
|
| 709 |
+
"zero-centered RMSNorm",
|
| 710 |
+
"tied input-output embeddings",
|
| 711 |
+
"optional multi-token prediction auxiliaries",
|
| 712 |
+
],
|
| 713 |
+
}
|
tiny_gdn/nn_common.py
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import torch
|
| 4 |
+
from torch import nn
|
| 5 |
+
from torch.nn.attention import SDPBackend
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
class RMSNorm(nn.Module):
|
| 9 |
+
"""Zero-centered RMSNorm as used by Qwen3-Next."""
|
| 10 |
+
|
| 11 |
+
def __init__(self, hidden_size: int, eps: float) -> None:
|
| 12 |
+
super().__init__()
|
| 13 |
+
self.weight = nn.Parameter(torch.zeros(hidden_size))
|
| 14 |
+
self.eps = eps
|
| 15 |
+
|
| 16 |
+
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
| 17 |
+
input_dtype = hidden_states.dtype
|
| 18 |
+
normalized = hidden_states.float()
|
| 19 |
+
normalized = normalized * torch.rsqrt(
|
| 20 |
+
normalized.square().mean(dim=-1, keepdim=True) + self.eps
|
| 21 |
+
)
|
| 22 |
+
normalized = normalized * (1.0 + self.weight.float())
|
| 23 |
+
return normalized.to(dtype=input_dtype)
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def attention_sdpa_backends(device: torch.device) -> list[SDPBackend]:
|
| 27 |
+
if device.type != "cuda":
|
| 28 |
+
return [SDPBackend.MATH]
|
| 29 |
+
return [
|
| 30 |
+
SDPBackend.CUDNN_ATTENTION,
|
| 31 |
+
SDPBackend.FLASH_ATTENTION,
|
| 32 |
+
SDPBackend.EFFICIENT_ATTENTION,
|
| 33 |
+
]
|
tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
tokenizer_config.json
ADDED
|
@@ -0,0 +1,74 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"add_bos_token": false,
|
| 3 |
+
"add_eos_token": false,
|
| 4 |
+
"bos_token": "<|begin_of_text|>",
|
| 5 |
+
"eos_token": "<|end_of_text|>",
|
| 6 |
+
"pad_token": "<|padding|>",
|
| 7 |
+
"unk_token": "<|unknown|>",
|
| 8 |
+
"model_max_length": 2048,
|
| 9 |
+
"clean_up_tokenization_spaces": false,
|
| 10 |
+
"tokenizer_class": "PreTrainedTokenizerFast",
|
| 11 |
+
"chat_template": "{%- set ns = namespace(xml_tools=none, tools_emitted=false) -%}\n{%- if xml_tools is defined and xml_tools -%}\n{%- set ns.xml_tools = xml_tools -%}\n{%- elif tools is defined and tools -%}\n{%- set ns.xml_tools = tools -%}\n{%- endif -%}\n{%- set tools_preamble = 'You may call one or more functions to assist with the user query.\\nYou are provided with function signatures within <tools></tools> XML tags:\\n<tools>\\n' -%}\n{%- set tools_epilogue = '</tools>\\n\\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\\n<tool_call>\\n{\"name\": <function-name>, \"arguments\": <args-json-object>}\\n</tool_call>' -%}\n{%- for message in messages -%}\n{%- if loop.first -%}{{- bos_token -}}{%- endif -%}\n{%- if loop.first and ns.xml_tools and message['role'] != 'system' -%}\n{{- '<|im_start|>system\\n' + tools_preamble -}}\n{%- for tool in ns.xml_tools -%}\n{%- if tool is string -%}{{- (tool | replace('<tools>', '') | replace('</tools>', '') | trim) + '\\n' -}}\n{%- else -%}{{- tool | tojson + '\\n' -}}\n{%- endif -%}\n{%- endfor -%}\n{{- tools_epilogue + '<|im_end|>\\n' -}}\n{%- set ns.tools_emitted = true -%}\n{%- endif -%}\n{%- set raw_role = message['role'] -%}\n{%- set role = 'user' if raw_role == 'tool' or raw_role == 'function' else raw_role -%}\n{%- set content_text = namespace(value='') -%}\n{%- if message['content'] is string -%}\n{%- set content_text.value = message['content'] -%}\n{%- elif message['content'] is iterable -%}\n{%- for item in message['content'] -%}\n{%- if item['type'] == 'text' -%}{%- set content_text.value = content_text.value + item['text'] -%}{%- endif -%}\n{%- endfor -%}\n{%- endif -%}\n{%- set is_tool = raw_role == 'tool' or raw_role == 'function' -%}\n{%- set prev_is_tool = loop.previtem is defined and (loop.previtem['role'] == 'tool' or loop.previtem['role'] == 'function') -%}\n{%- set next_is_tool = loop.nextitem is defined and (loop.nextitem['role'] == 'tool' or loop.nextitem['role'] == 'function') -%}\n{%- if raw_role == 'system' and not (content_text.value | trim) and not (ns.xml_tools and not ns.tools_emitted) -%}\n{%- else -%}\n{%- if is_tool and prev_is_tool -%}\n{{- '\\n' -}}\n{%- else -%}\n{{- '<|im_start|>' + role + '\\n' -}}\n{%- endif -%}\n{%- if raw_role == 'assistant' and enable_thinking is defined -%}\n{{- ('<|think|>\\n' if enable_thinking else '<|no_think|>\\n') -}}\n{%- endif -%}\n{%- if is_tool -%}\n{{- '<|tool_response|>\\n' + (content_text.value | replace('<tool_response>', '') | replace('</tool_response>', '') | trim) -}}\n{%- elif raw_role == 'system' -%}\n{{- content_text.value | replace('/system_override', '') | replace('/no_think', '') | replace('/think', '') | trim -}}\n{%- else -%}\n{{- content_text.value -}}\n{%- endif -%}\n{%- if raw_role == 'system' and ns.xml_tools and not ns.tools_emitted and '<tools>' not in content_text.value -%}\n{%- if content_text.value | trim -%}{{- '\\n\\n' -}}{%- endif -%}\n{{- tools_preamble -}}\n{%- for tool in ns.xml_tools -%}\n{%- if tool is string -%}{{- (tool | replace('<tools>', '') | replace('</tools>', '') | trim) + '\\n' -}}\n{%- else -%}{{- tool | tojson + '\\n' -}}\n{%- endif -%}\n{%- endfor -%}\n{{- tools_epilogue -}}\n{%- set ns.tools_emitted = true -%}\n{%- endif -%}\n{%- if raw_role == 'assistant' and message['tool_calls'] is defined and message['tool_calls'] -%}\n{%- if '<tool_call>' not in content_text.value -%}\n{%- for tool_call in message['tool_calls'] -%}\n{%- set fn = tool_call['function'] if tool_call['function'] is defined else tool_call -%}\n{%- if loop.first and not (content_text.value | trim) -%}\n{{- '<tool_call>\\n{\"name\": \"' + fn['name'] + '\", \"arguments\": ' -}}\n{%- else -%}\n{{- '\\n<tool_call>\\n{\"name\": \"' + fn['name'] + '\", \"arguments\": ' -}}\n{%- endif -%}\n{%- if fn['arguments'] is string -%}{{- fn['arguments'] -}}\n{%- else -%}{{- fn['arguments'] | tojson -}}\n{%- endif -%}\n{{- '}\\n</tool_call>' -}}\n{%- endfor -%}\n{%- endif -%}\n{%- endif -%}\n{%- if not (is_tool and next_is_tool) -%}\n{{- '<|im_end|>\\n' -}}\n{%- endif -%}\n{%- endif -%}\n{%- endfor -%}\n{%- if add_generation_prompt -%}\n{{- '<|im_start|>assistant\\n' -}}\n{%- if enable_thinking is defined -%}{{- ('<|think|>\\n' if enable_thinking else '<|no_think|>\\n') -}}{%- endif -%}\n{%- endif -%}",
|
| 12 |
+
"extra_special_tokens": [
|
| 13 |
+
"<|im_start|>",
|
| 14 |
+
"<|im_end|>",
|
| 15 |
+
"<|tool_call|>",
|
| 16 |
+
"<|tool_response|>",
|
| 17 |
+
"<think>",
|
| 18 |
+
"</think>",
|
| 19 |
+
"<|fim_prefix|>",
|
| 20 |
+
"<|fim_middle|>",
|
| 21 |
+
"<|fim_suffix|>",
|
| 22 |
+
"<|fim_pad|>",
|
| 23 |
+
"<|no_think|>",
|
| 24 |
+
"<|think|>",
|
| 25 |
+
"<|reserved_002|>",
|
| 26 |
+
"<|reserved_003|>",
|
| 27 |
+
"<|reserved_004|>",
|
| 28 |
+
"<|reserved_005|>",
|
| 29 |
+
"<|reserved_006|>",
|
| 30 |
+
"<|reserved_007|>",
|
| 31 |
+
"<|reserved_008|>",
|
| 32 |
+
"<|reserved_009|>",
|
| 33 |
+
"<|reserved_010|>",
|
| 34 |
+
"<|reserved_011|>",
|
| 35 |
+
"<|reserved_012|>",
|
| 36 |
+
"<|reserved_013|>",
|
| 37 |
+
"<|reserved_014|>",
|
| 38 |
+
"<|reserved_015|>",
|
| 39 |
+
"<|reserved_016|>",
|
| 40 |
+
"<|reserved_017|>",
|
| 41 |
+
"<|reserved_018|>",
|
| 42 |
+
"<|reserved_019|>",
|
| 43 |
+
"<|reserved_020|>",
|
| 44 |
+
"<|reserved_021|>",
|
| 45 |
+
"<|reserved_022|>",
|
| 46 |
+
"<|reserved_023|>",
|
| 47 |
+
"<|reserved_024|>",
|
| 48 |
+
"<|reserved_025|>",
|
| 49 |
+
"<|reserved_026|>",
|
| 50 |
+
"<|reserved_027|>",
|
| 51 |
+
"<|reserved_028|>",
|
| 52 |
+
"<|reserved_029|>",
|
| 53 |
+
"<|reserved_030|>",
|
| 54 |
+
"<|reserved_031|>",
|
| 55 |
+
"<|reserved_032|>",
|
| 56 |
+
"<|reserved_033|>",
|
| 57 |
+
"<|reserved_034|>",
|
| 58 |
+
"<|reserved_035|>",
|
| 59 |
+
"<|reserved_036|>",
|
| 60 |
+
"<|reserved_037|>",
|
| 61 |
+
"<|reserved_038|>",
|
| 62 |
+
"<|reserved_039|>",
|
| 63 |
+
"<|reserved_040|>",
|
| 64 |
+
"<|reserved_041|>",
|
| 65 |
+
"<|reserved_042|>",
|
| 66 |
+
"<|reserved_043|>",
|
| 67 |
+
"<|reserved_044|>",
|
| 68 |
+
"<|reserved_045|>",
|
| 69 |
+
"<|reserved_046|>",
|
| 70 |
+
"<|reserved_047|>",
|
| 71 |
+
"<|reserved_048|>",
|
| 72 |
+
"<|reserved_049|>"
|
| 73 |
+
]
|
| 74 |
+
}
|
validation.json
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"checkpoint": "checkpoint-00008400",
|
| 3 |
+
"ema": {
|
| 4 |
+
"batches": 16,
|
| 5 |
+
"elapsed_seconds": 7.330645765992813,
|
| 6 |
+
"loss": 3.6904296875,
|
| 7 |
+
"perplexity": 40.06205742476025,
|
| 8 |
+
"tokens": 524288,
|
| 9 |
+
"tokens_per_second": 71520.02930385683
|
| 10 |
+
},
|
| 11 |
+
"normal": {
|
| 12 |
+
"batches": 16,
|
| 13 |
+
"elapsed_seconds": 7.311252611019881,
|
| 14 |
+
"loss": 3.736328125,
|
| 15 |
+
"perplexity": 41.943695056893915,
|
| 16 |
+
"tokens": 524288,
|
| 17 |
+
"tokens_per_second": 71709.73674329993
|
| 18 |
+
},
|
| 19 |
+
"optimizer_step": 8400,
|
| 20 |
+
"type": "validation"
|
| 21 |
+
}
|
vocab.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
windows_fla_patches/fla/__init__.py
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
|
| 2 |
+
#
|
| 3 |
+
# Keep package import light. Eagerly importing fla.layers pulls every Triton
|
| 4 |
+
# kernel and breaks on Windows + Triton 3.7.
|
| 5 |
+
|
| 6 |
+
from pkgutil import extend_path
|
| 7 |
+
|
| 8 |
+
__path__ = extend_path(__path__, __name__)
|
| 9 |
+
__version__ = "0.5.2"
|
| 10 |
+
__all__: list[str] = []
|
windows_fla_patches/fla/layers/__init__.py
ADDED
|
@@ -0,0 +1,63 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
|
| 2 |
+
#
|
| 3 |
+
# Lazy layer exports — avoid compiling every Triton kernel at import time.
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
import importlib
|
| 8 |
+
from typing import Any
|
| 9 |
+
|
| 10 |
+
_EXPORTS: dict[str, tuple[str, str]] = {
|
| 11 |
+
"ABCAttention": (".abc", "ABCAttention"),
|
| 12 |
+
"Attention": (".attn", "Attention"),
|
| 13 |
+
"BasedLinearAttention": (".based", "BasedLinearAttention"),
|
| 14 |
+
"BitAttention": (".bitattn", "BitAttention"),
|
| 15 |
+
"Comba": (".comba", "Comba"),
|
| 16 |
+
"DeltaNet": (".delta_net", "DeltaNet"),
|
| 17 |
+
"DeltaFormerAttention": (".deltaformer", "DeltaFormerAttention"),
|
| 18 |
+
"ForgettingAttention": (".forgetting_attn", "ForgettingAttention"),
|
| 19 |
+
"GatedDeltaNet": (".gated_deltanet", "GatedDeltaNet"),
|
| 20 |
+
"GatedDeltaProduct": (".gated_deltaproduct", "GatedDeltaProduct"),
|
| 21 |
+
"GatedDeltaNet2": (".gdn2", "GatedDeltaNet2"),
|
| 22 |
+
"GatedLinearAttention": (".gla", "GatedLinearAttention"),
|
| 23 |
+
"GatedSlotAttention": (".gsa", "GatedSlotAttention"),
|
| 24 |
+
"HGRNAttention": (".hgrn", "HGRNAttention"),
|
| 25 |
+
"HGRN2Attention": (".hgrn2", "HGRN2Attention"),
|
| 26 |
+
"KimiDeltaAttention": (".kda", "KimiDeltaAttention"),
|
| 27 |
+
"LightNetAttention": (".lightnet", "LightNetAttention"),
|
| 28 |
+
"LinearAttention": (".linear_attn", "LinearAttention"),
|
| 29 |
+
"LogLinearMamba2": (".log_linear_mamba2", "LogLinearMamba2"),
|
| 30 |
+
"Mamba": (".mamba", "Mamba"),
|
| 31 |
+
"Mamba2": (".mamba2", "Mamba2"),
|
| 32 |
+
"Mamba3": (".mamba3", "Mamba3"),
|
| 33 |
+
"MesaNet": (".mesa_net", "MesaNet"),
|
| 34 |
+
"MultiheadLatentAttention": (".mla", "MultiheadLatentAttention"),
|
| 35 |
+
"MoBA": (".moba", "MoBA"),
|
| 36 |
+
"MomAttention": (".mom", "MomAttention"),
|
| 37 |
+
"MultiScaleRetention": (".multiscale_retention", "MultiScaleRetention"),
|
| 38 |
+
"NativeSparseAttention": (".nsa", "NativeSparseAttention"),
|
| 39 |
+
"Parallax": (".parallax", "Parallax"),
|
| 40 |
+
"PaTHAttention": (".path_attn", "PaTHAttention"),
|
| 41 |
+
"Raven": (".raven", "Raven"),
|
| 42 |
+
"ReBasedLinearAttention": (".rebased", "ReBasedLinearAttention"),
|
| 43 |
+
"RodimusAttention": (".rodimus", "RodimusAttention"),
|
| 44 |
+
"SlidingWindowSharedKeyAttention": (".rodimus", "SlidingWindowSharedKeyAttention"),
|
| 45 |
+
"RWKV6Attention": (".rwkv6", "RWKV6Attention"),
|
| 46 |
+
"RWKV7Attention": (".rwkv7", "RWKV7Attention"),
|
| 47 |
+
"WallAttention": (".wall_attn", "WallAttention"),
|
| 48 |
+
"YOCOCrossAttention": (".yoco", "YOCOCrossAttention"),
|
| 49 |
+
"YOCOGatedRetention": (".yoco", "YOCOGatedRetention"),
|
| 50 |
+
"YOCOSharedKVBuilder": (".yoco", "YOCOSharedKVBuilder"),
|
| 51 |
+
}
|
| 52 |
+
|
| 53 |
+
__all__ = list(_EXPORTS)
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def __getattr__(name: str) -> Any:
|
| 57 |
+
spec = _EXPORTS.get(name)
|
| 58 |
+
if spec is None:
|
| 59 |
+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
| 60 |
+
module_name, attr = spec
|
| 61 |
+
value = getattr(importlib.import_module(module_name, __name__), attr)
|
| 62 |
+
globals()[name] = value
|
| 63 |
+
return value
|
windows_fla_patches/fla/ops/__init__.py
ADDED
|
@@ -0,0 +1,76 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
|
| 2 |
+
#
|
| 3 |
+
# Lazy public exports so importing fla.ops.utils / fla.ops.gdn2 does not
|
| 4 |
+
# eagerly compile every Triton kernel (needed on Windows + Triton 3.7).
|
| 5 |
+
|
| 6 |
+
from __future__ import annotations
|
| 7 |
+
|
| 8 |
+
import importlib
|
| 9 |
+
from typing import Any
|
| 10 |
+
|
| 11 |
+
_EXPORTS: dict[str, str] = {
|
| 12 |
+
"chunk_abc": "fla.ops.abc",
|
| 13 |
+
"parallel_attn": "fla.ops.attn",
|
| 14 |
+
"fused_attnres": "fla.ops.attnres",
|
| 15 |
+
"fused_chunk_based": "fla.ops.based",
|
| 16 |
+
"parallel_based": "fla.ops.based",
|
| 17 |
+
"chunk_comba": "fla.ops.comba",
|
| 18 |
+
"fused_recurrent_comba": "fla.ops.comba",
|
| 19 |
+
"chunk_delta_rule": "fla.ops.delta_rule",
|
| 20 |
+
"fused_chunk_delta_rule": "fla.ops.delta_rule",
|
| 21 |
+
"fused_recurrent_delta_rule": "fla.ops.delta_rule",
|
| 22 |
+
"parallel_forgetting_attn": "fla.ops.forgetting_attn",
|
| 23 |
+
"chunk_gated_delta_rule": "fla.ops.gated_delta_rule",
|
| 24 |
+
"chunk_gdn": "fla.ops.gated_delta_rule",
|
| 25 |
+
"fused_recurrent_gated_delta_rule": "fla.ops.gated_delta_rule",
|
| 26 |
+
"fused_recurrent_gdn": "fla.ops.gated_delta_rule",
|
| 27 |
+
"chunk_dplr_delta_rule": "fla.ops.generalized_delta_rule",
|
| 28 |
+
"chunk_iplr_delta_rule": "fla.ops.generalized_delta_rule",
|
| 29 |
+
"fused_recurrent_dplr_delta_rule": "fla.ops.generalized_delta_rule",
|
| 30 |
+
"fused_recurrent_iplr_delta_rule": "fla.ops.generalized_delta_rule",
|
| 31 |
+
"chunk_gla": "fla.ops.gla",
|
| 32 |
+
"fused_chunk_gla": "fla.ops.gla",
|
| 33 |
+
"fused_recurrent_gla": "fla.ops.gla",
|
| 34 |
+
"chunk_gsa": "fla.ops.gsa",
|
| 35 |
+
"fused_recurrent_gsa": "fla.ops.gsa",
|
| 36 |
+
"fused_recurrent_hgrn": "fla.ops.hgrn",
|
| 37 |
+
"chunk_kda": "fla.ops.kda",
|
| 38 |
+
"fused_recurrent_kda": "fla.ops.kda",
|
| 39 |
+
"chunk_lightning_attn": "fla.ops.lightning_attn",
|
| 40 |
+
"fused_recurrent_lightning_attn": "fla.ops.lightning_attn",
|
| 41 |
+
"chunk_linear_attn": "fla.ops.linear_attn",
|
| 42 |
+
"fused_chunk_linear_attn": "fla.ops.linear_attn",
|
| 43 |
+
"fused_recurrent_linear_attn": "fla.ops.linear_attn",
|
| 44 |
+
"chunk_log_linear_attn": "fla.ops.log_linear_attn",
|
| 45 |
+
"chunk_mesa_net": "fla.ops.mesa_net",
|
| 46 |
+
"parallel_nsa": "fla.ops.nsa",
|
| 47 |
+
"parallel_parallax": "fla.ops.parallax",
|
| 48 |
+
"parallel_path_attn": "fla.ops.path_attn",
|
| 49 |
+
"chunk_retention": "fla.ops.retention",
|
| 50 |
+
"fused_chunk_retention": "fla.ops.retention",
|
| 51 |
+
"fused_recurrent_retention": "fla.ops.retention",
|
| 52 |
+
"parallel_retention": "fla.ops.retention",
|
| 53 |
+
"chunk_rwkv6": "fla.ops.rwkv6",
|
| 54 |
+
"fused_recurrent_rwkv6": "fla.ops.rwkv6",
|
| 55 |
+
"chunk_rwkv7": "fla.ops.rwkv7",
|
| 56 |
+
"fused_recurrent_rwkv7": "fla.ops.rwkv7",
|
| 57 |
+
"chunk_simple_gla": "fla.ops.simple_gla",
|
| 58 |
+
"fused_chunk_simple_gla": "fla.ops.simple_gla",
|
| 59 |
+
"fused_recurrent_simple_gla": "fla.ops.simple_gla",
|
| 60 |
+
"parallel_simple_gla": "fla.ops.simple_gla",
|
| 61 |
+
"parallel_wall_attn": "fla.ops.wall_attn",
|
| 62 |
+
"parallel_wall_attn_decode": "fla.ops.wall_attn",
|
| 63 |
+
"chunk_gdn2": "fla.ops.gdn2",
|
| 64 |
+
"fused_recurrent_gdn2": "fla.ops.gdn2",
|
| 65 |
+
}
|
| 66 |
+
|
| 67 |
+
__all__ = list(_EXPORTS)
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
def __getattr__(name: str) -> Any:
|
| 71 |
+
module_name = _EXPORTS.get(name)
|
| 72 |
+
if module_name is None:
|
| 73 |
+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
| 74 |
+
value = getattr(importlib.import_module(module_name), name)
|
| 75 |
+
globals()[name] = value
|
| 76 |
+
return value
|
windows_fla_patches/fla/ops/simple_gla/__init__.py
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
|
| 2 |
+
#
|
| 3 |
+
# This source code is licensed under the MIT license found in the
|
| 4 |
+
# LICENSE file in the root directory of this source tree.
|
| 5 |
+
# For a list of all contributors, visit:
|
| 6 |
+
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
|
| 7 |
+
|
| 8 |
+
from .chunk import chunk_simple_gla
|
| 9 |
+
from .fused_chunk import fused_chunk_simple_gla
|
| 10 |
+
from .fused_recurrent import fused_recurrent_simple_gla
|
| 11 |
+
|
| 12 |
+
# Triton 3.7 on Windows can fail while decorating parallel kernels at import time.
|
| 13 |
+
try:
|
| 14 |
+
from .parallel import parallel_simple_gla
|
| 15 |
+
except Exception: # noqa: BLE001
|
| 16 |
+
parallel_simple_gla = None
|
| 17 |
+
|
| 18 |
+
__all__ = [
|
| 19 |
+
'chunk_simple_gla',
|
| 20 |
+
'fused_chunk_simple_gla',
|
| 21 |
+
'fused_recurrent_simple_gla',
|
| 22 |
+
'parallel_simple_gla',
|
| 23 |
+
]
|