Upload Monostich-2 SFT EMA (smoltalk+hermes 1 epoch, step 10619, 21.70h)
Browse files- README.md +175 -0
- chat_template.jinja +17 -0
- config.json +82 -0
- inference.py +473 -0
- merges.txt +0 -0
- model.safetensors +3 -0
- requirements.txt +8 -0
- sft_config.json +37 -0
- special_token_ids.json +12 -0
- special_tokens_map.json +68 -0
- tiny_gdn/__init__.py +8 -0
- tiny_gdn/config.py +131 -0
- tiny_gdn/model.py +587 -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,175 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: apache-2.0
|
| 3 |
+
language:
|
| 4 |
+
- en
|
| 5 |
+
tags:
|
| 6 |
+
- text-generation
|
| 7 |
+
- causal-lm
|
| 8 |
+
- pytorch
|
| 9 |
+
- sft
|
| 10 |
+
- instruction-tuned
|
| 11 |
+
- chat
|
| 12 |
+
- hybrid
|
| 13 |
+
- gated-deltanet
|
| 14 |
+
- gqa
|
| 15 |
+
- monostich
|
| 16 |
+
pipeline_tag: text-generation
|
| 17 |
+
library_name: tiny_gdn
|
| 18 |
+
datasets:
|
| 19 |
+
- HuggingFaceTB/smoltalk
|
| 20 |
+
- NousResearch/Hermes-3-Dataset
|
| 21 |
+
base_model: kerzgrr/Monostich-2-base
|
| 22 |
+
model-index:
|
| 23 |
+
- name: Monostich-2
|
| 24 |
+
results: []
|
| 25 |
+
---
|
| 26 |
+
|
| 27 |
+
<div align="center">
|
| 28 |
+
|
| 29 |
+
# Monostich-2
|
| 30 |
+
|
| 31 |
+
### Instruction-tuned chat model (~150M) — Monostich-2 family
|
| 32 |
+
|
| 33 |
+
[](.)
|
| 34 |
+
[-green.svg)](.)
|
| 35 |
+
[](LICENSE)
|
| 36 |
+
[](https://huggingface.co/kerzgrr/Monostich-2-base)
|
| 37 |
+
|
| 38 |
+
*Second-generation Monostich-class model — hybrid GDN-2 + GQA, chat-tuned*
|
| 39 |
+
|
| 40 |
+
</div>
|
| 41 |
+
|
| 42 |
+
---
|
| 43 |
+
|
| 44 |
+
## What this is
|
| 45 |
+
|
| 46 |
+
**Monostich-2** is the **supervised fine-tuned (SFT) chat checkpoint** for the Monostich-2 family.
|
| 47 |
+
|
| 48 |
+
- Base (pretrain): [`kerzgrr/Monostich-2-base`](https://huggingface.co/kerzgrr/Monostich-2-base)
|
| 49 |
+
- Successor to [`kerzgrr/Monostich`](https://huggingface.co/kerzgrr/Monostich) (~100M LLaMA-style)
|
| 50 |
+
- Architecture: hybrid **Gated DeltaNet-2** + **gated GQA** (not plain LLaMA)
|
| 51 |
+
- This repo is **chat / instruction** (ChatML) — use the base repo for raw continuation
|
| 52 |
+
|
| 53 |
+
---
|
| 54 |
+
|
| 55 |
+
## Training
|
| 56 |
+
|
| 57 |
+
### Pretrain → SFT
|
| 58 |
+
|
| 59 |
+
| Stage | Details |
|
| 60 |
+
|-------|---------|
|
| 61 |
+
| **Base** | FineWeb-Edu pretrain → [`Monostich-2-base`](https://huggingface.co/kerzgrr/Monostich-2-base) |
|
| 62 |
+
| **SFT mix** | [HuggingFaceTB/smoltalk](https://huggingface.co/datasets/HuggingFaceTB/smoltalk) + [NousResearch/Hermes-3-Dataset](https://huggingface.co/datasets/NousResearch/Hermes-3-Dataset) |
|
| 63 |
+
| **Epochs** | 1 full epoch |
|
| 64 |
+
| **Wall time** | **21.70 hours** |
|
| 65 |
+
| **Final step** | optimizer step **10619** |
|
| 66 |
+
| **Weights** | EMA (Hub `model.safetensors` is EMA @ bfloat16) |
|
| 67 |
+
| **Seq length** | 8192 (packed SFT) |
|
| 68 |
+
| **Peak LR** | 1 × 10⁻⁴ AdamW, cosine → 10% min |
|
| 69 |
+
| **Final val loss (EMA)** | **1.532** (ppl ≈ 4.63) |
|
| 70 |
+
|
| 71 |
+
### Chat template (ChatML)
|
| 72 |
+
|
| 73 |
+
```
|
| 74 |
+
<|begin_of_text|><|im_start|>system
|
| 75 |
+
{system}<|im_end|>
|
| 76 |
+
<|im_start|>user
|
| 77 |
+
{user}<|im_end|>
|
| 78 |
+
<|im_start|>assistant
|
| 79 |
+
{assistant}<|im_end|>
|
| 80 |
+
```
|
| 81 |
+
|
| 82 |
+
Generation prompt ends at `<|im_start|>assistant\n`.
|
| 83 |
+
|
| 84 |
+
---
|
| 85 |
+
|
| 86 |
+
## Model Architecture
|
| 87 |
+
|
| 88 |
+
Same TinyGDN hybrid as the base (~149.1M params):
|
| 89 |
+
|
| 90 |
+
| | |
|
| 91 |
+
|--|--|
|
| 92 |
+
| **Layers** | 32 (GDN-2 ×3 + GQA every 4th) |
|
| 93 |
+
| **Hidden** | 512 |
|
| 94 |
+
| **MLP** | SwiGLU 1472 |
|
| 95 |
+
| **Attention** | 4 Q / 1 KV, head dim 128, partial RoPE |
|
| 96 |
+
| **Linear** | Gated DeltaNet-2, 4 heads × 128 |
|
| 97 |
+
| **Vocab** | 49,152 BPE |
|
| 98 |
+
|
| 99 |
+
---
|
| 100 |
+
|
| 101 |
+
## Install & run
|
| 102 |
+
|
| 103 |
+
```bash
|
| 104 |
+
pip install torch safetensors tokenizers huggingface_hub
|
| 105 |
+
hf download kerzgrr/Monostich-2 inference.py --local-dir .
|
| 106 |
+
python inference.py --prompt "What is the capital of France?"
|
| 107 |
+
```
|
| 108 |
+
|
| 109 |
+
`inference.py` auto-downloads weights/tokenizer/`tiny_gdn/` and **auto-installs** pinned `flash-linear-attention` (Windows applies Hub patches). Git required on `PATH`.
|
| 110 |
+
|
| 111 |
+
**Interactive chat:**
|
| 112 |
+
|
| 113 |
+
```bash
|
| 114 |
+
python inference.py
|
| 115 |
+
```
|
| 116 |
+
|
| 117 |
+
| Flag | Default | Description |
|
| 118 |
+
|------|---------|-------------|
|
| 119 |
+
| `--prompt` | — | One-shot user message |
|
| 120 |
+
| `--system` | — | Optional system prompt |
|
| 121 |
+
| `--temperature` | `0.7` | Sampling temperature |
|
| 122 |
+
| `--top-p` | `0.9` | Nucleus sampling |
|
| 123 |
+
| `--top-k` | `50` | Top-k |
|
| 124 |
+
| `--max-new-tokens` | `256` | Max generation length |
|
| 125 |
+
| `--device` | `cuda` if available | `cuda` / `cpu` |
|
| 126 |
+
|
| 127 |
+
---
|
| 128 |
+
|
| 129 |
+
## Limitations
|
| 130 |
+
|
| 131 |
+
- **~150M** research / edge chat model — not frontier quality
|
| 132 |
+
- Can **hallucinate**, repeat, fail simple arithmetic, lose multi-turn facts
|
| 133 |
+
- **Safety is incomplete** — do not deploy without your own filters
|
| 134 |
+
- Requires `flash-linear-attention` (GDN-2); **not GGUF / llama.cpp** compatible today
|
| 135 |
+
|
| 136 |
+
---
|
| 137 |
+
|
| 138 |
+
## Model family
|
| 139 |
+
|
| 140 |
+
| Model | Stage | Hub |
|
| 141 |
+
|-------|-------|-----|
|
| 142 |
+
| Monostich | SFT (~100M LLaMA) | [`kerzgrr/Monostich`](https://huggingface.co/kerzgrr/Monostich) |
|
| 143 |
+
| Monostich-2-base | Pretrain (~150M hybrid) | [`kerzgrr/Monostich-2-base`](https://huggingface.co/kerzgrr/Monostich-2-base) |
|
| 144 |
+
| **Monostich-2** | **SFT (~150M hybrid)** | **this repo** |
|
| 145 |
+
|
| 146 |
+
---
|
| 147 |
+
|
| 148 |
+
## Citation
|
| 149 |
+
|
| 150 |
+
```bibtex
|
| 151 |
+
@misc{monostich22026,
|
| 152 |
+
title={Monostich-2: A Hybrid GDN-2 + GQA Chat Model},
|
| 153 |
+
author={kerzgrr},
|
| 154 |
+
year={2026},
|
| 155 |
+
url={https://huggingface.co/kerzgrr/Monostich-2}
|
| 156 |
+
}
|
| 157 |
+
```
|
| 158 |
+
|
| 159 |
+
---
|
| 160 |
+
|
| 161 |
+
## Acknowledgments
|
| 162 |
+
|
| 163 |
+
- [flash-linear-attention](https://github.com/fla-org/flash-linear-attention) (Gated DeltaNet-2)
|
| 164 |
+
- [HuggingFaceTB/smoltalk](https://huggingface.co/datasets/HuggingFaceTB/smoltalk)
|
| 165 |
+
- [NousResearch/Hermes-3-Dataset](https://huggingface.co/datasets/NousResearch/Hermes-3-Dataset)
|
| 166 |
+
- Base: [`kerzgrr/Monostich-2-base`](https://huggingface.co/kerzgrr/Monostich-2-base)
|
| 167 |
+
|
| 168 |
+
---
|
| 169 |
+
|
| 170 |
+
<div align="center">
|
| 171 |
+
|
| 172 |
+
*A monostich is a poem of a single line — small, but complete.*
|
| 173 |
+
*Monostich-2 renews the form.*
|
| 174 |
+
|
| 175 |
+
</div>
|
chat_template.jinja
ADDED
|
@@ -0,0 +1,17 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{%- for message in messages -%}
|
| 2 |
+
{%- if loop.first -%}{{- bos_token -}}{%- endif -%}
|
| 3 |
+
{{- '<|im_start|>' + message['role'] + '\n' -}}
|
| 4 |
+
{%- if message['role'] == 'tool' -%}{{- '<|tool_response|>\n' -}}{%- endif -%}
|
| 5 |
+
{%- if message['content'] is string -%}
|
| 6 |
+
{{- message['content'] -}}
|
| 7 |
+
{%- elif message['content'] is iterable -%}
|
| 8 |
+
{%- for item in message['content'] -%}
|
| 9 |
+
{%- if item['type'] == 'text' -%}{{- item['text'] -}}{%- endif -%}
|
| 10 |
+
{%- endfor -%}
|
| 11 |
+
{%- endif -%}
|
| 12 |
+
{%- if message['tool_calls'] is defined and message['tool_calls'] -%}
|
| 13 |
+
{{- '\n<|tool_call|>\n' + (message['tool_calls'] | tojson) -}}
|
| 14 |
+
{%- endif -%}
|
| 15 |
+
{{- '<|im_end|>\n' -}}
|
| 16 |
+
{%- endfor -%}
|
| 17 |
+
{%- if add_generation_prompt -%}{{- '<|im_start|>assistant\n' -}}{%- endif -%}
|
config.json
ADDED
|
@@ -0,0 +1,82 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"allow_negative_eigenvalues": false,
|
| 3 |
+
"architecture": "TinyGDNForCausalLM",
|
| 4 |
+
"architectures": [
|
| 5 |
+
"TinyGDNForCausalLM"
|
| 6 |
+
],
|
| 7 |
+
"attention_dropout": 0.0,
|
| 8 |
+
"attention_head_dim": 128,
|
| 9 |
+
"base_model": "kerzgrr/Monostich-2-base",
|
| 10 |
+
"bos_token_id": 0,
|
| 11 |
+
"checkpoint_step": 10619,
|
| 12 |
+
"eos_token_id": 1,
|
| 13 |
+
"full_attention_interval": 4,
|
| 14 |
+
"hidden_size": 512,
|
| 15 |
+
"initializer_range": 0.02,
|
| 16 |
+
"intermediate_size": 1472,
|
| 17 |
+
"layer_types": [
|
| 18 |
+
"gdn2",
|
| 19 |
+
"gdn2",
|
| 20 |
+
"gdn2",
|
| 21 |
+
"full_attention",
|
| 22 |
+
"gdn2",
|
| 23 |
+
"gdn2",
|
| 24 |
+
"gdn2",
|
| 25 |
+
"full_attention",
|
| 26 |
+
"gdn2",
|
| 27 |
+
"gdn2",
|
| 28 |
+
"gdn2",
|
| 29 |
+
"full_attention",
|
| 30 |
+
"gdn2",
|
| 31 |
+
"gdn2",
|
| 32 |
+
"gdn2",
|
| 33 |
+
"full_attention",
|
| 34 |
+
"gdn2",
|
| 35 |
+
"gdn2",
|
| 36 |
+
"gdn2",
|
| 37 |
+
"full_attention",
|
| 38 |
+
"gdn2",
|
| 39 |
+
"gdn2",
|
| 40 |
+
"gdn2",
|
| 41 |
+
"full_attention",
|
| 42 |
+
"gdn2",
|
| 43 |
+
"gdn2",
|
| 44 |
+
"gdn2",
|
| 45 |
+
"full_attention",
|
| 46 |
+
"gdn2",
|
| 47 |
+
"gdn2",
|
| 48 |
+
"gdn2",
|
| 49 |
+
"full_attention"
|
| 50 |
+
],
|
| 51 |
+
"linear_conv_kernel_dim": 4,
|
| 52 |
+
"linear_expand_v": 1.0,
|
| 53 |
+
"linear_head_dim": 128,
|
| 54 |
+
"linear_num_heads": 4,
|
| 55 |
+
"linear_num_value_heads": 4,
|
| 56 |
+
"max_position_embeddings": 32768,
|
| 57 |
+
"model_family": "Monostich-2",
|
| 58 |
+
"model_name": "Monostich-2",
|
| 59 |
+
"model_type": "tiny_gdn",
|
| 60 |
+
"mtp_adapter_rank": 128,
|
| 61 |
+
"mtp_loss_weight": 0.0,
|
| 62 |
+
"mtp_num_heads": 0,
|
| 63 |
+
"num_attention_heads": 4,
|
| 64 |
+
"num_hidden_layers": 32,
|
| 65 |
+
"num_key_value_heads": 1,
|
| 66 |
+
"pad_token_id": 2,
|
| 67 |
+
"partial_rotary_factor": 0.5,
|
| 68 |
+
"rms_norm_eps": 1e-06,
|
| 69 |
+
"rope_theta": 1000000.0,
|
| 70 |
+
"sft_hours": 21.7,
|
| 71 |
+
"sft_val_loss_ema": 1.5321617420986058,
|
| 72 |
+
"sft_val_ppl_ema": 4.628170928002891,
|
| 73 |
+
"shared_layer_indices": [],
|
| 74 |
+
"stage": "sft",
|
| 75 |
+
"tie_word_embeddings": true,
|
| 76 |
+
"torch_dtype": "bfloat16",
|
| 77 |
+
"training_sequence_length": 8192,
|
| 78 |
+
"transformers_version": "4.45.0",
|
| 79 |
+
"unk_token_id": 3,
|
| 80 |
+
"vocab_size": 49152,
|
| 81 |
+
"weights": "ema"
|
| 82 |
+
}
|
inference.py
ADDED
|
@@ -0,0 +1,473 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
"""Standalone chat inference for kerzgrr/Monostich-2 (SFT).
|
| 3 |
+
|
| 4 |
+
Downloads model assets from the Hub (cached after first run), auto-installs
|
| 5 |
+
flash-linear-attention when needed, and streams ChatML assistant replies.
|
| 6 |
+
|
| 7 |
+
Examples:
|
| 8 |
+
python inference.py --prompt "What is the capital of France?"
|
| 9 |
+
python inference.py
|
| 10 |
+
python inference.py --prompt "Write a haiku about GPUs" --temperature 0.7
|
| 11 |
+
"""
|
| 12 |
+
|
| 13 |
+
from __future__ import annotations
|
| 14 |
+
|
| 15 |
+
import argparse
|
| 16 |
+
import json
|
| 17 |
+
import os
|
| 18 |
+
import platform
|
| 19 |
+
import shutil
|
| 20 |
+
import subprocess
|
| 21 |
+
import sys
|
| 22 |
+
import time
|
| 23 |
+
import warnings
|
| 24 |
+
from pathlib import Path
|
| 25 |
+
|
| 26 |
+
import torch
|
| 27 |
+
from safetensors.torch import load_file
|
| 28 |
+
from tokenizers import Tokenizer
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def _silence_runtime_warnings() -> None:
|
| 32 |
+
patterns = (
|
| 33 |
+
r"tl\.make_block_ptr is deprecated",
|
| 34 |
+
r"Memory efficient kernel not used because",
|
| 35 |
+
r"Memory Efficient attention has been runtime disabled",
|
| 36 |
+
r"Flash attention kernel not used because",
|
| 37 |
+
r"Torch was not compiled with flash attention",
|
| 38 |
+
r"cuDNN attention kernel not used because",
|
| 39 |
+
r"cuDNN attention has been runtime disabled",
|
| 40 |
+
)
|
| 41 |
+
for pattern in patterns:
|
| 42 |
+
warnings.filterwarnings("ignore", message=pattern)
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
REPO_ID = "kerzgrr/Monostich-2"
|
| 46 |
+
FLA_COMMIT = "cbb0a72efb55c18ca0ef4f298298317573ad2cb3"
|
| 47 |
+
FLA_REPO = "https://github.com/fla-org/flash-linear-attention.git"
|
| 48 |
+
PATCH_FILES = (
|
| 49 |
+
"fla/__init__.py",
|
| 50 |
+
"fla/ops/__init__.py",
|
| 51 |
+
"fla/layers/__init__.py",
|
| 52 |
+
"fla/ops/simple_gla/__init__.py",
|
| 53 |
+
)
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def _download(filename: str, local_dir: Path | None) -> Path:
|
| 57 |
+
from huggingface_hub import hf_hub_download
|
| 58 |
+
|
| 59 |
+
return Path(
|
| 60 |
+
hf_hub_download(
|
| 61 |
+
repo_id=REPO_ID,
|
| 62 |
+
filename=filename,
|
| 63 |
+
local_dir=str(local_dir) if local_dir else None,
|
| 64 |
+
)
|
| 65 |
+
)
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
def _run(cmd: list[str], *, cwd: Path | None = None, env: dict | None = None) -> None:
|
| 69 |
+
print("+", " ".join(cmd), flush=True)
|
| 70 |
+
merged = os.environ.copy()
|
| 71 |
+
if env:
|
| 72 |
+
merged.update(env)
|
| 73 |
+
merged.setdefault("PYTHONUTF8", "1")
|
| 74 |
+
merged.setdefault("PYTHONIOENCODING", "utf-8")
|
| 75 |
+
subprocess.check_call(cmd, cwd=str(cwd) if cwd else None, env=merged)
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
def _pip_install(*args: str) -> None:
|
| 79 |
+
_run([sys.executable, "-m", "pip", "install", *args])
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
def _fla_importable() -> tuple[bool, str]:
|
| 83 |
+
try:
|
| 84 |
+
from fla.layers.gdn2 import GatedDeltaNet2 # noqa: F401
|
| 85 |
+
except Exception as error: # noqa: BLE001
|
| 86 |
+
return False, str(error)
|
| 87 |
+
return True, ""
|
| 88 |
+
|
| 89 |
+
|
| 90 |
+
def _cache_root() -> Path:
|
| 91 |
+
override = os.environ.get("MONOSTICH_CACHE")
|
| 92 |
+
if override:
|
| 93 |
+
path = Path(override).expanduser().resolve()
|
| 94 |
+
else:
|
| 95 |
+
path = Path.home() / ".cache" / "monostich-2"
|
| 96 |
+
path.mkdir(parents=True, exist_ok=True)
|
| 97 |
+
return path
|
| 98 |
+
|
| 99 |
+
|
| 100 |
+
def _ensure_git() -> None:
|
| 101 |
+
if shutil.which("git") is None:
|
| 102 |
+
raise RuntimeError(
|
| 103 |
+
"git is required to auto-install flash-linear-attention. "
|
| 104 |
+
"Install Git and ensure it is on PATH."
|
| 105 |
+
)
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
def _apply_windows_fla_patches(fla_root: Path, local_dir: Path | None) -> None:
|
| 109 |
+
print("Applying Windows FLA import patches from the Hub …", flush=True)
|
| 110 |
+
for relative in PATCH_FILES:
|
| 111 |
+
source = _download(f"windows_fla_patches/{relative}", local_dir)
|
| 112 |
+
target = fla_root / relative
|
| 113 |
+
target.parent.mkdir(parents=True, exist_ok=True)
|
| 114 |
+
shutil.copy2(source, target)
|
| 115 |
+
print(f" patched {relative}", flush=True)
|
| 116 |
+
|
| 117 |
+
|
| 118 |
+
def _install_fla(local_dir: Path | None) -> None:
|
| 119 |
+
print("flash-linear-attention missing/broken — installing automatically …", flush=True)
|
| 120 |
+
_pip_install("einops", "numpy")
|
| 121 |
+
if platform.system() != "Windows":
|
| 122 |
+
_pip_install("--no-deps", f"git+{FLA_REPO}@{FLA_COMMIT}")
|
| 123 |
+
return
|
| 124 |
+
|
| 125 |
+
_ensure_git()
|
| 126 |
+
fla_root = _cache_root() / "flash-linear-attention"
|
| 127 |
+
if (fla_root / ".git").is_dir():
|
| 128 |
+
_run(["git", "fetch", "--depth", "1", "origin", FLA_COMMIT], cwd=fla_root)
|
| 129 |
+
_run(["git", "checkout", "--force", FLA_COMMIT], cwd=fla_root)
|
| 130 |
+
else:
|
| 131 |
+
if fla_root.exists():
|
| 132 |
+
shutil.rmtree(fla_root)
|
| 133 |
+
_run(["git", "clone", "--filter=blob:none", FLA_REPO, str(fla_root)])
|
| 134 |
+
_run(["git", "fetch", "--depth", "1", "origin", FLA_COMMIT], cwd=fla_root)
|
| 135 |
+
_run(["git", "checkout", "--force", FLA_COMMIT], cwd=fla_root)
|
| 136 |
+
|
| 137 |
+
_apply_windows_fla_patches(fla_root, local_dir)
|
| 138 |
+
_pip_install("--no-build-isolation", "--no-deps", "-e", str(fla_root))
|
| 139 |
+
|
| 140 |
+
|
| 141 |
+
def _ensure_fla(local_dir: Path | None) -> None:
|
| 142 |
+
ok, error = _fla_importable()
|
| 143 |
+
if ok:
|
| 144 |
+
return
|
| 145 |
+
print(f"FLA not ready ({error})", flush=True)
|
| 146 |
+
try:
|
| 147 |
+
_install_fla(local_dir)
|
| 148 |
+
except Exception as install_error: # noqa: BLE001
|
| 149 |
+
raise RuntimeError(
|
| 150 |
+
"Automatic flash-linear-attention install failed.\n"
|
| 151 |
+
f"Original import error: {error}\n"
|
| 152 |
+
f"Install error: {install_error}"
|
| 153 |
+
) from install_error
|
| 154 |
+
|
| 155 |
+
for name in list(sys.modules):
|
| 156 |
+
if name == "fla" or name.startswith("fla."):
|
| 157 |
+
del sys.modules[name]
|
| 158 |
+
|
| 159 |
+
ok, error = _fla_importable()
|
| 160 |
+
if not ok:
|
| 161 |
+
raise RuntimeError(
|
| 162 |
+
"flash-linear-attention installed but still failed to import "
|
| 163 |
+
f"GatedDeltaNet2: {error}"
|
| 164 |
+
)
|
| 165 |
+
print("flash-linear-attention ready.", flush=True)
|
| 166 |
+
|
| 167 |
+
|
| 168 |
+
def _ensure_tiny_gdn(local_dir: Path | None) -> Path:
|
| 169 |
+
here = Path(__file__).resolve().parent
|
| 170 |
+
if (here / "tiny_gdn" / "__init__.py").is_file():
|
| 171 |
+
return here
|
| 172 |
+
if local_dir and (local_dir / "tiny_gdn" / "__init__.py").is_file():
|
| 173 |
+
return local_dir
|
| 174 |
+
for name in ("tiny_gdn/__init__.py", "tiny_gdn/config.py", "tiny_gdn/model.py"):
|
| 175 |
+
_download(name, local_dir)
|
| 176 |
+
return _download("tiny_gdn/__init__.py", local_dir).parent.parent
|
| 177 |
+
|
| 178 |
+
|
| 179 |
+
def _sample(
|
| 180 |
+
logits: torch.Tensor,
|
| 181 |
+
*,
|
| 182 |
+
temperature: float,
|
| 183 |
+
top_p: float,
|
| 184 |
+
top_k: int,
|
| 185 |
+
generator: torch.Generator,
|
| 186 |
+
) -> int:
|
| 187 |
+
logits = logits.float()
|
| 188 |
+
if temperature <= 1e-5:
|
| 189 |
+
return int(torch.argmax(logits).item())
|
| 190 |
+
logits = logits / temperature
|
| 191 |
+
if 0 < top_k < logits.shape[-1]:
|
| 192 |
+
threshold = torch.topk(logits, top_k).values[-1]
|
| 193 |
+
logits = logits.masked_fill(logits < threshold, -torch.inf)
|
| 194 |
+
if top_p < 1.0:
|
| 195 |
+
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
|
| 196 |
+
probs = torch.softmax(sorted_logits, dim=-1)
|
| 197 |
+
remove = torch.cumsum(probs, dim=-1) > top_p
|
| 198 |
+
remove[1:] = remove[:-1].clone()
|
| 199 |
+
remove[0] = False
|
| 200 |
+
sorted_logits = sorted_logits.masked_fill(remove, -torch.inf)
|
| 201 |
+
logits = torch.full_like(logits, -torch.inf)
|
| 202 |
+
logits.scatter_(0, sorted_indices, sorted_logits)
|
| 203 |
+
probs = torch.softmax(logits, dim=-1)
|
| 204 |
+
return int(torch.multinomial(probs, 1, generator=generator).item())
|
| 205 |
+
|
| 206 |
+
|
| 207 |
+
def _apply_repetition_penalty(
|
| 208 |
+
logits: torch.Tensor,
|
| 209 |
+
token_ids: list[int],
|
| 210 |
+
penalty: float,
|
| 211 |
+
window: int,
|
| 212 |
+
) -> torch.Tensor:
|
| 213 |
+
if penalty == 1.0 or not token_ids:
|
| 214 |
+
return logits
|
| 215 |
+
recent = token_ids[-window:] if window > 0 else token_ids
|
| 216 |
+
unique = torch.tensor(list(set(recent)), dtype=torch.long, device=logits.device)
|
| 217 |
+
score = logits[unique]
|
| 218 |
+
logits[unique] = torch.where(score > 0, score / penalty, score * penalty)
|
| 219 |
+
return logits
|
| 220 |
+
|
| 221 |
+
|
| 222 |
+
def encode_chat(tokenizer: Tokenizer, messages: list[dict[str, str]]) -> list[int]:
|
| 223 |
+
bos = tokenizer.token_to_id("<|begin_of_text|>")
|
| 224 |
+
im_start = tokenizer.token_to_id("<|im_start|>")
|
| 225 |
+
im_end = tokenizer.token_to_id("<|im_end|>")
|
| 226 |
+
if bos is None or im_start is None or im_end is None:
|
| 227 |
+
raise RuntimeError("Tokenizer missing ChatML specials")
|
| 228 |
+
ids = [bos]
|
| 229 |
+
newline = tokenizer.encode("\n", add_special_tokens=False).ids
|
| 230 |
+
for message in messages:
|
| 231 |
+
role = message["role"]
|
| 232 |
+
content = message["content"]
|
| 233 |
+
ids.append(im_start)
|
| 234 |
+
ids.extend(tokenizer.encode(f"{role}\n", add_special_tokens=False).ids)
|
| 235 |
+
ids.extend(tokenizer.encode(content, add_special_tokens=False).ids)
|
| 236 |
+
ids.append(im_end)
|
| 237 |
+
ids.extend(newline)
|
| 238 |
+
ids.append(im_start)
|
| 239 |
+
ids.extend(tokenizer.encode("assistant\n", add_special_tokens=False).ids)
|
| 240 |
+
return ids
|
| 241 |
+
|
| 242 |
+
|
| 243 |
+
@torch.inference_mode()
|
| 244 |
+
def generate(
|
| 245 |
+
model,
|
| 246 |
+
tokenizer: Tokenizer,
|
| 247 |
+
prompt_ids: list[int],
|
| 248 |
+
*,
|
| 249 |
+
max_new_tokens: int,
|
| 250 |
+
context_length: int,
|
| 251 |
+
temperature: float,
|
| 252 |
+
top_p: float,
|
| 253 |
+
top_k: int,
|
| 254 |
+
repetition_penalty: float,
|
| 255 |
+
repetition_window: int,
|
| 256 |
+
seed: int,
|
| 257 |
+
stream: bool,
|
| 258 |
+
device: torch.device,
|
| 259 |
+
) -> tuple[str, int, str]:
|
| 260 |
+
eos_id = int(model.config.eos_token_id)
|
| 261 |
+
im_end = tokenizer.token_to_id("<|im_end|>")
|
| 262 |
+
stop = {eos_id}
|
| 263 |
+
if im_end is not None:
|
| 264 |
+
stop.add(im_end)
|
| 265 |
+
|
| 266 |
+
token_ids = list(prompt_ids[-context_length:])
|
| 267 |
+
generated: list[int] = []
|
| 268 |
+
decoded = ""
|
| 269 |
+
stop_reason = "max_new_tokens"
|
| 270 |
+
generator = torch.Generator(device=device)
|
| 271 |
+
generator.manual_seed(seed)
|
| 272 |
+
started = time.perf_counter()
|
| 273 |
+
|
| 274 |
+
for _ in range(max_new_tokens):
|
| 275 |
+
context = token_ids[-context_length:]
|
| 276 |
+
input_ids = torch.tensor([context], dtype=torch.long, device=device)
|
| 277 |
+
output = model(input_ids, return_logits=True, logits_to_keep=1)
|
| 278 |
+
if output.logits is None:
|
| 279 |
+
raise RuntimeError("Model returned no logits")
|
| 280 |
+
next_logits = _apply_repetition_penalty(
|
| 281 |
+
output.logits[0, -1],
|
| 282 |
+
token_ids,
|
| 283 |
+
repetition_penalty,
|
| 284 |
+
repetition_window,
|
| 285 |
+
)
|
| 286 |
+
next_id = _sample(
|
| 287 |
+
next_logits,
|
| 288 |
+
temperature=temperature,
|
| 289 |
+
top_p=top_p,
|
| 290 |
+
top_k=top_k,
|
| 291 |
+
generator=generator,
|
| 292 |
+
)
|
| 293 |
+
if next_id in stop:
|
| 294 |
+
stop_reason = "stop"
|
| 295 |
+
break
|
| 296 |
+
token_ids.append(next_id)
|
| 297 |
+
generated.append(next_id)
|
| 298 |
+
current = tokenizer.decode(generated, skip_special_tokens=True)
|
| 299 |
+
delta = current[len(decoded) :] if current.startswith(decoded) else current
|
| 300 |
+
decoded = current
|
| 301 |
+
if stream and delta:
|
| 302 |
+
print(delta, end="", flush=True)
|
| 303 |
+
|
| 304 |
+
if stream:
|
| 305 |
+
print(flush=True)
|
| 306 |
+
elapsed = time.perf_counter() - started
|
| 307 |
+
tps = len(generated) / max(elapsed, 1e-9)
|
| 308 |
+
if stream:
|
| 309 |
+
print(
|
| 310 |
+
f"[done] tokens={len(generated)} stop={stop_reason} {tps:.1f} tok/s",
|
| 311 |
+
file=sys.stderr,
|
| 312 |
+
flush=True,
|
| 313 |
+
)
|
| 314 |
+
return decoded, len(generated), stop_reason
|
| 315 |
+
|
| 316 |
+
|
| 317 |
+
def parse_args() -> argparse.Namespace:
|
| 318 |
+
parser = argparse.ArgumentParser(description="Monostich-2 SFT chat inference")
|
| 319 |
+
parser.add_argument("--prompt", default=None, help="Single user prompt")
|
| 320 |
+
parser.add_argument("--system", default="", help="Optional system prompt")
|
| 321 |
+
parser.add_argument("--max-new-tokens", type=int, default=256)
|
| 322 |
+
parser.add_argument("--temperature", type=float, default=0.7)
|
| 323 |
+
parser.add_argument("--top-p", type=float, default=0.9)
|
| 324 |
+
parser.add_argument("--top-k", type=int, default=50)
|
| 325 |
+
parser.add_argument("--repetition-penalty", type=float, default=1.08)
|
| 326 |
+
parser.add_argument("--repetition-window", type=int, default=256)
|
| 327 |
+
parser.add_argument("--context-length", type=int, default=2048)
|
| 328 |
+
parser.add_argument("--seed", type=int, default=42)
|
| 329 |
+
parser.add_argument(
|
| 330 |
+
"--device",
|
| 331 |
+
default="cuda" if torch.cuda.is_available() else "cpu",
|
| 332 |
+
choices=["cuda", "cpu"],
|
| 333 |
+
)
|
| 334 |
+
parser.add_argument("--no-stream", action="store_true")
|
| 335 |
+
parser.add_argument("--local-dir", default=None)
|
| 336 |
+
parser.add_argument("--repo-id", default=REPO_ID)
|
| 337 |
+
return parser.parse_args()
|
| 338 |
+
|
| 339 |
+
|
| 340 |
+
def main() -> int:
|
| 341 |
+
_silence_runtime_warnings()
|
| 342 |
+
args = parse_args()
|
| 343 |
+
global REPO_ID
|
| 344 |
+
REPO_ID = args.repo_id
|
| 345 |
+
local_dir = Path(args.local_dir).resolve() if args.local_dir else None
|
| 346 |
+
|
| 347 |
+
print(f"Loading Monostich-2 from huggingface.co/{REPO_ID} …", flush=True)
|
| 348 |
+
try:
|
| 349 |
+
package_root = _ensure_tiny_gdn(local_dir)
|
| 350 |
+
except Exception as error: # noqa: BLE001
|
| 351 |
+
print(f"Failed to resolve tiny_gdn package: {error}", file=sys.stderr)
|
| 352 |
+
return 1
|
| 353 |
+
|
| 354 |
+
if str(package_root) not in sys.path:
|
| 355 |
+
sys.path.insert(0, str(package_root))
|
| 356 |
+
|
| 357 |
+
try:
|
| 358 |
+
_ensure_fla(local_dir)
|
| 359 |
+
except Exception as error: # noqa: BLE001
|
| 360 |
+
print(str(error), file=sys.stderr)
|
| 361 |
+
return 1
|
| 362 |
+
|
| 363 |
+
try:
|
| 364 |
+
from tiny_gdn import TinyGDNConfig, TinyGDNForCausalLM
|
| 365 |
+
except ImportError as error:
|
| 366 |
+
print(f"Could not import tiny_gdn: {error}", file=sys.stderr)
|
| 367 |
+
return 1
|
| 368 |
+
|
| 369 |
+
weights_path = _download("model.safetensors", local_dir)
|
| 370 |
+
tok_path = _download("tokenizer.json", local_dir)
|
| 371 |
+
cfg_path = _download("config.json", local_dir)
|
| 372 |
+
|
| 373 |
+
raw = json.loads(cfg_path.read_text(encoding="utf-8"))
|
| 374 |
+
from dataclasses import fields
|
| 375 |
+
|
| 376 |
+
allowed = {item.name for item in fields(TinyGDNConfig)}
|
| 377 |
+
payload = {key: value for key, value in raw.items() if key in allowed}
|
| 378 |
+
if "shared_layer_indices" in payload:
|
| 379 |
+
payload["shared_layer_indices"] = tuple(payload["shared_layer_indices"])
|
| 380 |
+
config = TinyGDNConfig(**payload)
|
| 381 |
+
|
| 382 |
+
device = torch.device(args.device)
|
| 383 |
+
if device.type == "cuda" and not torch.cuda.is_available():
|
| 384 |
+
print("CUDA requested but unavailable; falling back to CPU.", flush=True)
|
| 385 |
+
device = torch.device("cpu")
|
| 386 |
+
dtype = torch.bfloat16 if device.type == "cuda" else torch.float32
|
| 387 |
+
|
| 388 |
+
print(
|
| 389 |
+
f"Building TinyGDN ({config.num_hidden_layers}L / {config.hidden_size}d) "
|
| 390 |
+
f"on {device} …",
|
| 391 |
+
flush=True,
|
| 392 |
+
)
|
| 393 |
+
model = TinyGDNForCausalLM(config)
|
| 394 |
+
state = load_file(str(weights_path), device="cpu")
|
| 395 |
+
model.load_state_dict(state, strict=True)
|
| 396 |
+
del state
|
| 397 |
+
model = model.to(device=device, dtype=dtype)
|
| 398 |
+
model.eval()
|
| 399 |
+
model.requires_grad_(False)
|
| 400 |
+
|
| 401 |
+
tokenizer = Tokenizer.from_file(str(tok_path))
|
| 402 |
+
context_length = min(args.context_length, config.max_position_embeddings)
|
| 403 |
+
stream = not args.no_stream
|
| 404 |
+
|
| 405 |
+
def run_chat(messages: list[dict[str, str]]) -> str:
|
| 406 |
+
prompt_ids = encode_chat(tokenizer, messages)
|
| 407 |
+
text, _, _ = generate(
|
| 408 |
+
model,
|
| 409 |
+
tokenizer,
|
| 410 |
+
prompt_ids,
|
| 411 |
+
max_new_tokens=args.max_new_tokens,
|
| 412 |
+
context_length=context_length,
|
| 413 |
+
temperature=args.temperature,
|
| 414 |
+
top_p=args.top_p,
|
| 415 |
+
top_k=args.top_k,
|
| 416 |
+
repetition_penalty=args.repetition_penalty,
|
| 417 |
+
repetition_window=args.repetition_window,
|
| 418 |
+
seed=args.seed,
|
| 419 |
+
stream=stream,
|
| 420 |
+
device=device,
|
| 421 |
+
)
|
| 422 |
+
return text
|
| 423 |
+
|
| 424 |
+
if args.prompt is not None:
|
| 425 |
+
messages: list[dict[str, str]] = []
|
| 426 |
+
if args.system.strip():
|
| 427 |
+
messages.append({"role": "system", "content": args.system.strip()})
|
| 428 |
+
messages.append({"role": "user", "content": args.prompt})
|
| 429 |
+
if stream:
|
| 430 |
+
print("assistant> ", end="", flush=True)
|
| 431 |
+
text = run_chat(messages)
|
| 432 |
+
if not stream:
|
| 433 |
+
print(text)
|
| 434 |
+
return 0
|
| 435 |
+
|
| 436 |
+
print(
|
| 437 |
+
"Interactive chat. Commands: /exit /quit /reset\n"
|
| 438 |
+
"This is the SFT chat model (ChatML).",
|
| 439 |
+
flush=True,
|
| 440 |
+
)
|
| 441 |
+
history: list[dict[str, str]] = []
|
| 442 |
+
if args.system.strip():
|
| 443 |
+
history.append({"role": "system", "content": args.system.strip()})
|
| 444 |
+
while True:
|
| 445 |
+
try:
|
| 446 |
+
user_input = input("user> ")
|
| 447 |
+
except (EOFError, KeyboardInterrupt):
|
| 448 |
+
print()
|
| 449 |
+
break
|
| 450 |
+
text = user_input.strip()
|
| 451 |
+
if not text:
|
| 452 |
+
continue
|
| 453 |
+
if text.lower() in {"/exit", "/quit"}:
|
| 454 |
+
break
|
| 455 |
+
if text.lower() == "/reset":
|
| 456 |
+
history = []
|
| 457 |
+
if args.system.strip():
|
| 458 |
+
history.append({"role": "system", "content": args.system.strip()})
|
| 459 |
+
print("(history cleared)", flush=True)
|
| 460 |
+
continue
|
| 461 |
+
turn = history + [{"role": "user", "content": text}]
|
| 462 |
+
if stream:
|
| 463 |
+
print("assistant> ", end="", flush=True)
|
| 464 |
+
reply = run_chat(turn)
|
| 465 |
+
history = turn + [{"role": "assistant", "content": reply}]
|
| 466 |
+
if not stream:
|
| 467 |
+
print(f"assistant> {reply}")
|
| 468 |
+
print(flush=True)
|
| 469 |
+
return 0
|
| 470 |
+
|
| 471 |
+
|
| 472 |
+
if __name__ == "__main__":
|
| 473 |
+
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:78a40abe31ebda87173b3dea67b8e89766ce36e1f313fab41f0f10d369a851d6
|
| 3 |
+
size 298279400
|
requirements.txt
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
+
# flash-linear-attention is auto-installed by inference.py on first run.
|
sft_config.json
ADDED
|
@@ -0,0 +1,37 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"adam_beta1": 0.9,
|
| 3 |
+
"adam_beta2": 0.95,
|
| 4 |
+
"adam_epsilon": 1e-08,
|
| 5 |
+
"cpu_prefetch_factor": 2,
|
| 6 |
+
"cpu_workers": 2,
|
| 7 |
+
"dataset_manifest": "data/tokenized/smoltalk-hermes-sft/manifest.json",
|
| 8 |
+
"ema_inv_gamma": 1.0,
|
| 9 |
+
"ema_max_decay": 0.9999,
|
| 10 |
+
"ema_power": 0.75,
|
| 11 |
+
"enable_gradient_checkpointing": true,
|
| 12 |
+
"full_shuffle": true,
|
| 13 |
+
"gradient_accumulation_steps": 16,
|
| 14 |
+
"gradient_clip_norm": 1.0,
|
| 15 |
+
"initial_run_dir": "runs/tiny-gdn-150m-smoltalk-sft",
|
| 16 |
+
"initial_weights": "ema",
|
| 17 |
+
"learning_rate": 0.0001,
|
| 18 |
+
"log_every_optimizer_steps": 1,
|
| 19 |
+
"logit_z_loss_coefficient": 0.0001,
|
| 20 |
+
"maximum_checkpoints": 5,
|
| 21 |
+
"maximum_optimizer_steps": null,
|
| 22 |
+
"micro_batch_size": 1,
|
| 23 |
+
"minimum_learning_rate_ratio": 0.1,
|
| 24 |
+
"muon_learning_rate": 0.01,
|
| 25 |
+
"muon_momentum": 0.95,
|
| 26 |
+
"optimizer": "adamw",
|
| 27 |
+
"output_dir": "runs/tiny-gdn-150m-smoltalk-hermes-sft",
|
| 28 |
+
"require_complete_pretraining": false,
|
| 29 |
+
"require_complete_sft_source": false,
|
| 30 |
+
"save_every_optimizer_steps": 100,
|
| 31 |
+
"seed": 20260720,
|
| 32 |
+
"sequence_length": 8192,
|
| 33 |
+
"shuffle_block_sequences": 512,
|
| 34 |
+
"validation_batches": 16,
|
| 35 |
+
"warmup_ratio": 0.01,
|
| 36 |
+
"weight_decay": 0.0
|
| 37 |
+
}
|
special_token_ids.json
ADDED
|
@@ -0,0 +1,12 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"bos_token": "<|begin_of_text|>",
|
| 3 |
+
"eos_token": "<|end_of_text|>",
|
| 4 |
+
"pad_token": "<|padding|>",
|
| 5 |
+
"unk_token": "<|unknown|>",
|
| 6 |
+
"im_start": "<|im_start|>",
|
| 7 |
+
"im_end": "<|im_end|>",
|
| 8 |
+
"bos_token_id": 0,
|
| 9 |
+
"eos_token_id": 1,
|
| 10 |
+
"pad_token_id": 2,
|
| 11 |
+
"unk_token_id": 3
|
| 12 |
+
}
|
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_start|>",
|
| 12 |
+
"<|think_end|>",
|
| 13 |
+
"<|fim_prefix|>",
|
| 14 |
+
"<|fim_middle|>",
|
| 15 |
+
"<|fim_suffix|>",
|
| 16 |
+
"<|fim_pad|>",
|
| 17 |
+
"<|reserved_000|>",
|
| 18 |
+
"<|reserved_001|>",
|
| 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,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from tiny_gdn.config import TinyGDNConfig
|
| 2 |
+
from tiny_gdn.model import TinyGDNForCausalLM, TinyGDNOutput
|
| 3 |
+
|
| 4 |
+
__all__ = [
|
| 5 |
+
"TinyGDNConfig",
|
| 6 |
+
"TinyGDNForCausalLM",
|
| 7 |
+
"TinyGDNOutput",
|
| 8 |
+
]
|
tiny_gdn/config.py
ADDED
|
@@ -0,0 +1,131 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import json
|
| 4 |
+
from dataclasses import asdict, dataclass
|
| 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 |
+
|
| 36 |
+
max_position_embeddings: int = 32_768
|
| 37 |
+
training_sequence_length: int = 2_048
|
| 38 |
+
rms_norm_eps: float = 1e-6
|
| 39 |
+
initializer_range: float = 0.02
|
| 40 |
+
tie_word_embeddings: bool = True
|
| 41 |
+
shared_layer_indices: tuple[int, ...] = ()
|
| 42 |
+
|
| 43 |
+
# MTP is an opt-in ablation at this scale; static MTP is not assumed to
|
| 44 |
+
# improve a 150M model without a controlled pilot.
|
| 45 |
+
mtp_num_heads: int = 0
|
| 46 |
+
mtp_adapter_rank: int = 128
|
| 47 |
+
mtp_loss_weight: float = 0.0
|
| 48 |
+
|
| 49 |
+
bos_token_id: int = 0
|
| 50 |
+
eos_token_id: int = 1
|
| 51 |
+
pad_token_id: int = 2
|
| 52 |
+
unk_token_id: int = 3
|
| 53 |
+
|
| 54 |
+
def __post_init__(self) -> None:
|
| 55 |
+
if self.vocab_size <= 0 or self.vocab_size > 65_536:
|
| 56 |
+
raise ValueError("vocab_size must fit the uint16 token dataset")
|
| 57 |
+
if self.hidden_size != self.num_attention_heads * self.attention_head_dim:
|
| 58 |
+
raise ValueError("hidden_size must equal num_attention_heads * attention_head_dim")
|
| 59 |
+
if self.hidden_size != self.linear_num_heads * self.linear_head_dim:
|
| 60 |
+
raise ValueError("hidden_size must equal linear_num_heads * linear_head_dim")
|
| 61 |
+
if self.linear_num_value_heads < self.linear_num_heads:
|
| 62 |
+
raise ValueError("linear_num_value_heads must be at least linear_num_heads")
|
| 63 |
+
if self.linear_num_value_heads % self.linear_num_heads != 0:
|
| 64 |
+
raise ValueError("linear_num_value_heads must be divisible by linear_num_heads")
|
| 65 |
+
if self.num_attention_heads % self.num_key_value_heads != 0:
|
| 66 |
+
raise ValueError("num_attention_heads must be divisible by num_key_value_heads")
|
| 67 |
+
if self.num_hidden_layers % self.full_attention_interval != 0:
|
| 68 |
+
raise ValueError("num_hidden_layers must be divisible by full_attention_interval")
|
| 69 |
+
if not 0.0 < self.partial_rotary_factor <= 1.0:
|
| 70 |
+
raise ValueError("partial_rotary_factor must be in (0, 1]")
|
| 71 |
+
rotary_dim = int(self.attention_head_dim * self.partial_rotary_factor)
|
| 72 |
+
if rotary_dim <= 0 or rotary_dim % 2:
|
| 73 |
+
raise ValueError("The partial rotary dimension must be positive and even")
|
| 74 |
+
if self.training_sequence_length > self.max_position_embeddings:
|
| 75 |
+
raise ValueError("training_sequence_length exceeds max_position_embeddings")
|
| 76 |
+
if len(set(self.shared_layer_indices)) != len(self.shared_layer_indices):
|
| 77 |
+
raise ValueError("shared_layer_indices must be unique")
|
| 78 |
+
if any(
|
| 79 |
+
index < 0 or index >= self.num_hidden_layers
|
| 80 |
+
for index in self.shared_layer_indices
|
| 81 |
+
):
|
| 82 |
+
raise ValueError("shared_layer_indices contains an invalid layer")
|
| 83 |
+
if self.mtp_num_heads < 0:
|
| 84 |
+
raise ValueError("mtp_num_heads cannot be negative")
|
| 85 |
+
if self.mtp_num_heads and self.mtp_adapter_rank <= 0:
|
| 86 |
+
raise ValueError("mtp_adapter_rank must be positive when MTP is enabled")
|
| 87 |
+
if not 0.0 <= self.mtp_loss_weight <= 1.0:
|
| 88 |
+
raise ValueError("mtp_loss_weight must be between zero and one")
|
| 89 |
+
for token_id in (
|
| 90 |
+
self.bos_token_id,
|
| 91 |
+
self.eos_token_id,
|
| 92 |
+
self.pad_token_id,
|
| 93 |
+
self.unk_token_id,
|
| 94 |
+
):
|
| 95 |
+
if not 0 <= token_id < self.vocab_size:
|
| 96 |
+
raise ValueError(f"Special token ID {token_id} is outside the vocabulary")
|
| 97 |
+
|
| 98 |
+
@property
|
| 99 |
+
def layer_types(self) -> tuple[str, ...]:
|
| 100 |
+
return tuple(
|
| 101 |
+
"full_attention" if (index + 1) % self.full_attention_interval == 0 else "gdn2"
|
| 102 |
+
for index in range(self.num_hidden_layers)
|
| 103 |
+
)
|
| 104 |
+
|
| 105 |
+
@property
|
| 106 |
+
def rotary_dim(self) -> int:
|
| 107 |
+
return int(self.attention_head_dim * self.partial_rotary_factor)
|
| 108 |
+
|
| 109 |
+
@property
|
| 110 |
+
def effective_num_layers(self) -> int:
|
| 111 |
+
return self.num_hidden_layers + len(self.shared_layer_indices)
|
| 112 |
+
|
| 113 |
+
def to_dict(self) -> dict[str, Any]:
|
| 114 |
+
payload = asdict(self)
|
| 115 |
+
payload["layer_types"] = list(self.layer_types)
|
| 116 |
+
return payload
|
| 117 |
+
|
| 118 |
+
def save_json(self, path: Path) -> None:
|
| 119 |
+
path.parent.mkdir(parents=True, exist_ok=True)
|
| 120 |
+
path.write_text(
|
| 121 |
+
json.dumps(self.to_dict(), indent=2, sort_keys=True) + "\n",
|
| 122 |
+
encoding="utf-8",
|
| 123 |
+
)
|
| 124 |
+
|
| 125 |
+
@classmethod
|
| 126 |
+
def from_json(cls, path: Path) -> TinyGDNConfig:
|
| 127 |
+
payload = json.loads(path.read_text(encoding="utf-8"))
|
| 128 |
+
payload.pop("layer_types", None)
|
| 129 |
+
if "shared_layer_indices" in payload:
|
| 130 |
+
payload["shared_layer_indices"] = tuple(payload["shared_layer_indices"])
|
| 131 |
+
return cls(**payload)
|
tiny_gdn/model.py
ADDED
|
@@ -0,0 +1,587 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 SDPBackend, sdpa_kernel
|
| 13 |
+
from torch.utils.checkpoint import checkpoint
|
| 14 |
+
|
| 15 |
+
from tiny_gdn.config import TinyGDNConfig
|
| 16 |
+
|
| 17 |
+
try:
|
| 18 |
+
# Import the module directly — `from fla.layers import GatedDeltaNet2`
|
| 19 |
+
# executes layers/__init__.py and eagerly loads every attention kernel.
|
| 20 |
+
from fla.layers.gdn2 import GatedDeltaNet2
|
| 21 |
+
except ImportError as import_error:
|
| 22 |
+
GatedDeltaNet2 = None
|
| 23 |
+
FLA_IMPORT_ERROR: ImportError | None = import_error
|
| 24 |
+
else:
|
| 25 |
+
FLA_IMPORT_ERROR = None
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
@dataclass
|
| 29 |
+
class TinyGDNOutput:
|
| 30 |
+
loss: torch.Tensor | None
|
| 31 |
+
logits: torch.Tensor | None
|
| 32 |
+
main_loss: torch.Tensor | None
|
| 33 |
+
mtp_loss: torch.Tensor | None
|
| 34 |
+
z_loss: torch.Tensor | None
|
| 35 |
+
hidden_states: torch.Tensor | None = None
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
class RMSNorm(nn.Module):
|
| 39 |
+
"""Zero-centered RMSNorm as used by Qwen3-Next."""
|
| 40 |
+
|
| 41 |
+
def __init__(self, hidden_size: int, eps: float) -> None:
|
| 42 |
+
super().__init__()
|
| 43 |
+
self.weight = nn.Parameter(torch.zeros(hidden_size))
|
| 44 |
+
self.eps = eps
|
| 45 |
+
|
| 46 |
+
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
| 47 |
+
input_dtype = hidden_states.dtype
|
| 48 |
+
normalized = hidden_states.float()
|
| 49 |
+
normalized = normalized * torch.rsqrt(normalized.square().mean(dim=-1, keepdim=True) + self.eps)
|
| 50 |
+
normalized = normalized * (1.0 + self.weight.float())
|
| 51 |
+
return normalized.to(dtype=input_dtype)
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
class RotaryEmbedding(nn.Module):
|
| 55 |
+
def __init__(self, rotary_dim: int, rope_theta: float) -> None:
|
| 56 |
+
super().__init__()
|
| 57 |
+
inverse_frequency = 1.0 / (
|
| 58 |
+
rope_theta
|
| 59 |
+
** (
|
| 60 |
+
torch.arange(0, rotary_dim, 2, dtype=torch.float32)
|
| 61 |
+
/ rotary_dim
|
| 62 |
+
)
|
| 63 |
+
)
|
| 64 |
+
self.rotary_dim = rotary_dim
|
| 65 |
+
self.register_buffer("inverse_frequency", inverse_frequency, persistent=False)
|
| 66 |
+
|
| 67 |
+
def forward(
|
| 68 |
+
self,
|
| 69 |
+
sequence_length: int,
|
| 70 |
+
device: torch.device,
|
| 71 |
+
dtype: torch.dtype,
|
| 72 |
+
position_offset: int = 0,
|
| 73 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 74 |
+
positions = torch.arange(
|
| 75 |
+
position_offset,
|
| 76 |
+
position_offset + sequence_length,
|
| 77 |
+
device=device,
|
| 78 |
+
dtype=torch.float32,
|
| 79 |
+
)
|
| 80 |
+
frequencies = torch.outer(positions, self.inverse_frequency.float())
|
| 81 |
+
embeddings = torch.cat((frequencies, frequencies), dim=-1)
|
| 82 |
+
return embeddings.cos().to(dtype=dtype), embeddings.sin().to(dtype=dtype)
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
def rotate_half(hidden_states: torch.Tensor) -> torch.Tensor:
|
| 86 |
+
first, second = hidden_states.chunk(2, dim=-1)
|
| 87 |
+
return torch.cat((-second, first), dim=-1)
|
| 88 |
+
|
| 89 |
+
|
| 90 |
+
def apply_rotary_embedding(
|
| 91 |
+
query: torch.Tensor,
|
| 92 |
+
key: torch.Tensor,
|
| 93 |
+
cosine: torch.Tensor,
|
| 94 |
+
sine: torch.Tensor,
|
| 95 |
+
rotary_dim: int,
|
| 96 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 97 |
+
cosine = cosine[None, None, :, :]
|
| 98 |
+
sine = sine[None, None, :, :]
|
| 99 |
+
query_rotary, query_pass = query[..., :rotary_dim], query[..., rotary_dim:]
|
| 100 |
+
key_rotary, key_pass = key[..., :rotary_dim], key[..., rotary_dim:]
|
| 101 |
+
query_rotary = query_rotary * cosine + rotate_half(query_rotary) * sine
|
| 102 |
+
key_rotary = key_rotary * cosine + rotate_half(key_rotary) * sine
|
| 103 |
+
return (
|
| 104 |
+
torch.cat((query_rotary, query_pass), dim=-1),
|
| 105 |
+
torch.cat((key_rotary, key_pass), dim=-1),
|
| 106 |
+
)
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
class GatedGroupedQueryAttention(nn.Module):
|
| 110 |
+
"""QK-normalized, partially rotary GQA with a learned sigmoid output gate."""
|
| 111 |
+
|
| 112 |
+
def __init__(self, config: TinyGDNConfig) -> None:
|
| 113 |
+
super().__init__()
|
| 114 |
+
self.num_heads = config.num_attention_heads
|
| 115 |
+
self.num_key_value_heads = config.num_key_value_heads
|
| 116 |
+
self.head_dim = config.attention_head_dim
|
| 117 |
+
self.rotary_dim = config.rotary_dim
|
| 118 |
+
self.dropout = config.attention_dropout
|
| 119 |
+
|
| 120 |
+
query_size = self.num_heads * self.head_dim
|
| 121 |
+
key_value_size = self.num_key_value_heads * self.head_dim
|
| 122 |
+
self.q_gate_proj = nn.Linear(config.hidden_size, query_size * 2, bias=False)
|
| 123 |
+
self.k_proj = nn.Linear(config.hidden_size, key_value_size, bias=False)
|
| 124 |
+
self.v_proj = nn.Linear(config.hidden_size, key_value_size, bias=False)
|
| 125 |
+
self.o_proj = nn.Linear(query_size, config.hidden_size, bias=False)
|
| 126 |
+
self.q_norm = RMSNorm(self.head_dim, config.rms_norm_eps)
|
| 127 |
+
self.k_norm = RMSNorm(self.head_dim, config.rms_norm_eps)
|
| 128 |
+
self.rotary = RotaryEmbedding(self.rotary_dim, config.rope_theta)
|
| 129 |
+
|
| 130 |
+
def _attention_mask(
|
| 131 |
+
self,
|
| 132 |
+
attention_mask: torch.Tensor | None,
|
| 133 |
+
sequence_length: int,
|
| 134 |
+
device: torch.device,
|
| 135 |
+
) -> torch.Tensor | None:
|
| 136 |
+
if attention_mask is None:
|
| 137 |
+
return None
|
| 138 |
+
if attention_mask.ndim != 2:
|
| 139 |
+
raise ValueError("attention_mask must have shape [batch, sequence]")
|
| 140 |
+
if attention_mask.shape[1] != sequence_length:
|
| 141 |
+
raise ValueError("attention_mask sequence length does not match input")
|
| 142 |
+
|
| 143 |
+
causal = torch.ones(
|
| 144 |
+
sequence_length,
|
| 145 |
+
sequence_length,
|
| 146 |
+
dtype=torch.bool,
|
| 147 |
+
device=device,
|
| 148 |
+
).tril()
|
| 149 |
+
valid_keys = attention_mask[:, None, None, :].to(dtype=torch.bool, device=device)
|
| 150 |
+
return causal[None, None, :, :] & valid_keys
|
| 151 |
+
|
| 152 |
+
def forward(
|
| 153 |
+
self,
|
| 154 |
+
hidden_states: torch.Tensor,
|
| 155 |
+
attention_mask: torch.Tensor | None = None,
|
| 156 |
+
) -> torch.Tensor:
|
| 157 |
+
batch_size, sequence_length, _ = hidden_states.shape
|
| 158 |
+
query_and_gate = self.q_gate_proj(hidden_states)
|
| 159 |
+
query, output_gate = query_and_gate.chunk(2, dim=-1)
|
| 160 |
+
|
| 161 |
+
query = query.view(batch_size, sequence_length, self.num_heads, self.head_dim)
|
| 162 |
+
key = self.k_proj(hidden_states).view(
|
| 163 |
+
batch_size,
|
| 164 |
+
sequence_length,
|
| 165 |
+
self.num_key_value_heads,
|
| 166 |
+
self.head_dim,
|
| 167 |
+
)
|
| 168 |
+
value = self.v_proj(hidden_states).view(
|
| 169 |
+
batch_size,
|
| 170 |
+
sequence_length,
|
| 171 |
+
self.num_key_value_heads,
|
| 172 |
+
self.head_dim,
|
| 173 |
+
)
|
| 174 |
+
|
| 175 |
+
query = self.q_norm(query).transpose(1, 2)
|
| 176 |
+
key = self.k_norm(key).transpose(1, 2)
|
| 177 |
+
value = value.transpose(1, 2)
|
| 178 |
+
|
| 179 |
+
cosine, sine = self.rotary(
|
| 180 |
+
sequence_length,
|
| 181 |
+
device=hidden_states.device,
|
| 182 |
+
dtype=query.dtype,
|
| 183 |
+
)
|
| 184 |
+
query, key = apply_rotary_embedding(
|
| 185 |
+
query,
|
| 186 |
+
key,
|
| 187 |
+
cosine,
|
| 188 |
+
sine,
|
| 189 |
+
rotary_dim=self.rotary_dim,
|
| 190 |
+
)
|
| 191 |
+
|
| 192 |
+
sdpa_mask = self._attention_mask(
|
| 193 |
+
attention_mask,
|
| 194 |
+
sequence_length,
|
| 195 |
+
hidden_states.device,
|
| 196 |
+
)
|
| 197 |
+
sdpa_options = {
|
| 198 |
+
"attn_mask": sdpa_mask,
|
| 199 |
+
"dropout_p": self.dropout if self.training else 0.0,
|
| 200 |
+
"is_causal": sdpa_mask is None,
|
| 201 |
+
"enable_gqa": True,
|
| 202 |
+
}
|
| 203 |
+
# Prefer Flash / mem-efficient when available; fall back to MATH for
|
| 204 |
+
# Windows PyTorch builds that ship without FlashAttention kernels.
|
| 205 |
+
backends = (
|
| 206 |
+
[
|
| 207 |
+
SDPBackend.FLASH_ATTENTION,
|
| 208 |
+
SDPBackend.EFFICIENT_ATTENTION,
|
| 209 |
+
SDPBackend.CUDNN_ATTENTION,
|
| 210 |
+
SDPBackend.MATH,
|
| 211 |
+
]
|
| 212 |
+
if query.is_cuda
|
| 213 |
+
else [SDPBackend.MATH]
|
| 214 |
+
)
|
| 215 |
+
with sdpa_kernel(backends):
|
| 216 |
+
attention_output = F.scaled_dot_product_attention(
|
| 217 |
+
query,
|
| 218 |
+
key,
|
| 219 |
+
value,
|
| 220 |
+
**sdpa_options,
|
| 221 |
+
)
|
| 222 |
+
attention_output = attention_output.transpose(1, 2).reshape(
|
| 223 |
+
batch_size,
|
| 224 |
+
sequence_length,
|
| 225 |
+
-1,
|
| 226 |
+
)
|
| 227 |
+
attention_output = attention_output * torch.sigmoid(output_gate)
|
| 228 |
+
return self.o_proj(attention_output)
|
| 229 |
+
|
| 230 |
+
|
| 231 |
+
class SwiGLU(nn.Module):
|
| 232 |
+
def __init__(self, config: TinyGDNConfig) -> None:
|
| 233 |
+
super().__init__()
|
| 234 |
+
self.gate_up_proj = nn.Linear(
|
| 235 |
+
config.hidden_size,
|
| 236 |
+
config.intermediate_size * 2,
|
| 237 |
+
bias=False,
|
| 238 |
+
)
|
| 239 |
+
self.down_proj = nn.Linear(
|
| 240 |
+
config.intermediate_size,
|
| 241 |
+
config.hidden_size,
|
| 242 |
+
bias=False,
|
| 243 |
+
)
|
| 244 |
+
|
| 245 |
+
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
| 246 |
+
gate, up = self.gate_up_proj(hidden_states).chunk(2, dim=-1)
|
| 247 |
+
return self.down_proj(F.silu(gate) * up)
|
| 248 |
+
|
| 249 |
+
|
| 250 |
+
class TinyGDNBlock(nn.Module):
|
| 251 |
+
def __init__(self, config: TinyGDNConfig, layer_index: int) -> None:
|
| 252 |
+
super().__init__()
|
| 253 |
+
layer_type = config.layer_types[layer_index]
|
| 254 |
+
self.layer_type = layer_type
|
| 255 |
+
self.token_mixer_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
| 256 |
+
self.mlp_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
| 257 |
+
|
| 258 |
+
if layer_type == "gdn2":
|
| 259 |
+
if GatedDeltaNet2 is None:
|
| 260 |
+
raise ImportError(
|
| 261 |
+
"Gated DeltaNet-2 requires the pinned flash-linear-attention dependency"
|
| 262 |
+
) from FLA_IMPORT_ERROR
|
| 263 |
+
self.token_mixer = GatedDeltaNet2(
|
| 264 |
+
hidden_size=config.hidden_size,
|
| 265 |
+
expand_v=config.linear_expand_v,
|
| 266 |
+
head_dim=config.linear_head_dim,
|
| 267 |
+
num_heads=config.linear_num_heads,
|
| 268 |
+
num_v_heads=config.linear_num_value_heads,
|
| 269 |
+
mode="chunk",
|
| 270 |
+
use_short_conv=True,
|
| 271 |
+
allow_neg_eigval=config.allow_negative_eigenvalues,
|
| 272 |
+
conv_size=config.linear_conv_kernel_dim,
|
| 273 |
+
conv_bias=False,
|
| 274 |
+
layer_idx=layer_index,
|
| 275 |
+
norm_eps=config.rms_norm_eps,
|
| 276 |
+
)
|
| 277 |
+
elif layer_type == "full_attention":
|
| 278 |
+
self.token_mixer = GatedGroupedQueryAttention(config)
|
| 279 |
+
else:
|
| 280 |
+
raise ValueError(f"Unsupported layer type: {layer_type}")
|
| 281 |
+
|
| 282 |
+
self.mlp = SwiGLU(config)
|
| 283 |
+
|
| 284 |
+
def forward(
|
| 285 |
+
self,
|
| 286 |
+
hidden_states: torch.Tensor,
|
| 287 |
+
attention_mask: torch.Tensor | None = None,
|
| 288 |
+
) -> torch.Tensor:
|
| 289 |
+
residual = hidden_states
|
| 290 |
+
normalized = self.token_mixer_norm(hidden_states)
|
| 291 |
+
if self.layer_type == "gdn2":
|
| 292 |
+
mixed, _, _ = self.token_mixer(
|
| 293 |
+
normalized,
|
| 294 |
+
attention_mask=attention_mask,
|
| 295 |
+
use_cache=False,
|
| 296 |
+
)
|
| 297 |
+
else:
|
| 298 |
+
mixed = self.token_mixer(normalized, attention_mask=attention_mask)
|
| 299 |
+
hidden_states = residual + mixed
|
| 300 |
+
hidden_states = hidden_states + self.mlp(self.mlp_norm(hidden_states))
|
| 301 |
+
return hidden_states
|
| 302 |
+
|
| 303 |
+
|
| 304 |
+
class MultiTokenPredictionAdapter(nn.Module):
|
| 305 |
+
"""A lightweight residual adapter for one additional prediction horizon."""
|
| 306 |
+
|
| 307 |
+
def __init__(self, config: TinyGDNConfig) -> None:
|
| 308 |
+
super().__init__()
|
| 309 |
+
self.norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
| 310 |
+
self.down_proj = nn.Linear(
|
| 311 |
+
config.hidden_size,
|
| 312 |
+
config.mtp_adapter_rank,
|
| 313 |
+
bias=False,
|
| 314 |
+
)
|
| 315 |
+
self.up_proj = nn.Linear(
|
| 316 |
+
config.mtp_adapter_rank,
|
| 317 |
+
config.hidden_size,
|
| 318 |
+
bias=False,
|
| 319 |
+
)
|
| 320 |
+
|
| 321 |
+
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
| 322 |
+
adapted = self.up_proj(F.silu(self.down_proj(self.norm(hidden_states))))
|
| 323 |
+
return hidden_states + adapted
|
| 324 |
+
|
| 325 |
+
|
| 326 |
+
class TinyGDNForCausalLM(nn.Module):
|
| 327 |
+
def __init__(self, config: TinyGDNConfig) -> None:
|
| 328 |
+
super().__init__()
|
| 329 |
+
self.config = config
|
| 330 |
+
self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size)
|
| 331 |
+
self.layers = nn.ModuleList(
|
| 332 |
+
TinyGDNBlock(config, layer_index)
|
| 333 |
+
for layer_index in range(config.num_hidden_layers)
|
| 334 |
+
)
|
| 335 |
+
self.final_norm = RMSNorm(config.hidden_size, config.rms_norm_eps)
|
| 336 |
+
self.mtp_adapters = nn.ModuleList(
|
| 337 |
+
MultiTokenPredictionAdapter(config)
|
| 338 |
+
for _ in range(config.mtp_num_heads)
|
| 339 |
+
)
|
| 340 |
+
self.gradient_checkpointing = False
|
| 341 |
+
|
| 342 |
+
self.apply(self._initialize_module)
|
| 343 |
+
self._initialize_residual_projections()
|
| 344 |
+
|
| 345 |
+
def _initialize_module(self, module: nn.Module) -> None:
|
| 346 |
+
if isinstance(module, nn.Linear):
|
| 347 |
+
nn.init.normal_(
|
| 348 |
+
module.weight,
|
| 349 |
+
mean=0.0,
|
| 350 |
+
std=self.config.initializer_range,
|
| 351 |
+
)
|
| 352 |
+
if module.bias is not None:
|
| 353 |
+
nn.init.zeros_(module.bias)
|
| 354 |
+
elif isinstance(module, nn.Embedding):
|
| 355 |
+
nn.init.normal_(
|
| 356 |
+
module.weight,
|
| 357 |
+
mean=0.0,
|
| 358 |
+
std=self.config.initializer_range,
|
| 359 |
+
)
|
| 360 |
+
|
| 361 |
+
def _initialize_residual_projections(self) -> None:
|
| 362 |
+
residual_std = self.config.initializer_range / math.sqrt(
|
| 363 |
+
2 * self.config.num_hidden_layers
|
| 364 |
+
)
|
| 365 |
+
for layer in self.layers:
|
| 366 |
+
nn.init.normal_(
|
| 367 |
+
layer.token_mixer.o_proj.weight,
|
| 368 |
+
mean=0.0,
|
| 369 |
+
std=residual_std,
|
| 370 |
+
)
|
| 371 |
+
nn.init.normal_(
|
| 372 |
+
layer.mlp.down_proj.weight,
|
| 373 |
+
mean=0.0,
|
| 374 |
+
std=residual_std,
|
| 375 |
+
)
|
| 376 |
+
for adapter in self.mtp_adapters:
|
| 377 |
+
nn.init.normal_(adapter.up_proj.weight, mean=0.0, std=residual_std)
|
| 378 |
+
|
| 379 |
+
def enable_gradient_checkpointing(self, enabled: bool = True) -> None:
|
| 380 |
+
self.gradient_checkpointing = enabled
|
| 381 |
+
|
| 382 |
+
def project_to_vocabulary(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
| 383 |
+
return F.linear(hidden_states, self.embed_tokens.weight)
|
| 384 |
+
|
| 385 |
+
def _run_layer(
|
| 386 |
+
self,
|
| 387 |
+
layer: TinyGDNBlock,
|
| 388 |
+
hidden_states: torch.Tensor,
|
| 389 |
+
attention_mask: torch.Tensor | None,
|
| 390 |
+
) -> torch.Tensor:
|
| 391 |
+
if self.gradient_checkpointing and self.training:
|
| 392 |
+
return checkpoint(
|
| 393 |
+
layer,
|
| 394 |
+
hidden_states,
|
| 395 |
+
attention_mask,
|
| 396 |
+
use_reentrant=False,
|
| 397 |
+
)
|
| 398 |
+
return layer(hidden_states, attention_mask)
|
| 399 |
+
|
| 400 |
+
def _causal_loss(
|
| 401 |
+
self,
|
| 402 |
+
hidden_states: torch.Tensor,
|
| 403 |
+
labels: torch.Tensor,
|
| 404 |
+
target_offset: int,
|
| 405 |
+
adapter: nn.Module | None = None,
|
| 406 |
+
compute_z_loss: bool = False,
|
| 407 |
+
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
| 408 |
+
if target_offset < 0:
|
| 409 |
+
raise ValueError("target_offset cannot be negative")
|
| 410 |
+
if target_offset and hidden_states.shape[1] <= target_offset:
|
| 411 |
+
raise ValueError(
|
| 412 |
+
f"Sequence length must exceed target offset {target_offset}"
|
| 413 |
+
)
|
| 414 |
+
if target_offset:
|
| 415 |
+
prediction_states = hidden_states[:, :-target_offset, :]
|
| 416 |
+
targets = labels[:, target_offset:].contiguous()
|
| 417 |
+
else:
|
| 418 |
+
prediction_states = hidden_states
|
| 419 |
+
targets = labels.contiguous()
|
| 420 |
+
if adapter is not None:
|
| 421 |
+
prediction_states = adapter(prediction_states)
|
| 422 |
+
logits = self.project_to_vocabulary(prediction_states)
|
| 423 |
+
cross_entropy = F.cross_entropy(
|
| 424 |
+
logits.reshape(-1, self.config.vocab_size),
|
| 425 |
+
targets.reshape(-1),
|
| 426 |
+
ignore_index=-100,
|
| 427 |
+
)
|
| 428 |
+
z_loss = None
|
| 429 |
+
if compute_z_loss:
|
| 430 |
+
valid_targets = targets.ne(-100)
|
| 431 |
+
log_partition = torch.logsumexp(logits.float(), dim=-1)
|
| 432 |
+
z_loss = log_partition.square()[valid_targets].mean()
|
| 433 |
+
return cross_entropy, z_loss
|
| 434 |
+
|
| 435 |
+
def forward(
|
| 436 |
+
self,
|
| 437 |
+
input_ids: torch.Tensor,
|
| 438 |
+
labels: torch.Tensor | None = None,
|
| 439 |
+
attention_mask: torch.Tensor | None = None,
|
| 440 |
+
*,
|
| 441 |
+
return_logits: bool = True,
|
| 442 |
+
return_hidden_states: bool = False,
|
| 443 |
+
labels_are_shifted: bool = False,
|
| 444 |
+
include_mtp_loss: bool = True,
|
| 445 |
+
mtp_loss_weight: float | None = None,
|
| 446 |
+
z_loss_coefficient: float = 0.0,
|
| 447 |
+
logits_to_keep: int | None = None,
|
| 448 |
+
) -> TinyGDNOutput:
|
| 449 |
+
if input_ids.ndim != 2:
|
| 450 |
+
raise ValueError("input_ids must have shape [batch, sequence]")
|
| 451 |
+
if input_ids.shape[1] > self.config.max_position_embeddings:
|
| 452 |
+
raise ValueError("Input exceeds max_position_embeddings")
|
| 453 |
+
if labels is not None and labels.shape != input_ids.shape:
|
| 454 |
+
raise ValueError("labels must have the same shape as input_ids")
|
| 455 |
+
if z_loss_coefficient < 0.0:
|
| 456 |
+
raise ValueError("z_loss_coefficient cannot be negative")
|
| 457 |
+
if logits_to_keep is not None and logits_to_keep <= 0:
|
| 458 |
+
raise ValueError("logits_to_keep must be positive")
|
| 459 |
+
effective_mtp_weight = (
|
| 460 |
+
self.config.mtp_loss_weight
|
| 461 |
+
if mtp_loss_weight is None
|
| 462 |
+
else mtp_loss_weight
|
| 463 |
+
)
|
| 464 |
+
if not 0.0 <= effective_mtp_weight <= 1.0:
|
| 465 |
+
raise ValueError("mtp_loss_weight must be between zero and one")
|
| 466 |
+
|
| 467 |
+
hidden_states = self.embed_tokens(input_ids)
|
| 468 |
+
shared_layer_indices = set(self.config.shared_layer_indices)
|
| 469 |
+
for layer_index, layer in enumerate(self.layers):
|
| 470 |
+
hidden_states = self._run_layer(
|
| 471 |
+
layer,
|
| 472 |
+
hidden_states,
|
| 473 |
+
attention_mask,
|
| 474 |
+
)
|
| 475 |
+
if layer_index in shared_layer_indices:
|
| 476 |
+
hidden_states = self._run_layer(
|
| 477 |
+
layer,
|
| 478 |
+
hidden_states,
|
| 479 |
+
attention_mask,
|
| 480 |
+
)
|
| 481 |
+
hidden_states = self.final_norm(hidden_states)
|
| 482 |
+
|
| 483 |
+
main_loss = None
|
| 484 |
+
mtp_loss = None
|
| 485 |
+
z_loss = None
|
| 486 |
+
total_loss = None
|
| 487 |
+
if labels is not None:
|
| 488 |
+
main_target_offset = 0 if labels_are_shifted else 1
|
| 489 |
+
main_loss, z_loss = self._causal_loss(
|
| 490 |
+
hidden_states,
|
| 491 |
+
labels,
|
| 492 |
+
target_offset=main_target_offset,
|
| 493 |
+
compute_z_loss=z_loss_coefficient > 0.0,
|
| 494 |
+
)
|
| 495 |
+
if self.mtp_adapters and include_mtp_loss:
|
| 496 |
+
auxiliary_losses = [
|
| 497 |
+
self._causal_loss(
|
| 498 |
+
hidden_states,
|
| 499 |
+
labels,
|
| 500 |
+
target_offset=(
|
| 501 |
+
head_index + 1
|
| 502 |
+
if labels_are_shifted
|
| 503 |
+
else head_index + 2
|
| 504 |
+
),
|
| 505 |
+
adapter=adapter,
|
| 506 |
+
)[0]
|
| 507 |
+
for head_index, adapter in enumerate(self.mtp_adapters)
|
| 508 |
+
]
|
| 509 |
+
mtp_loss = torch.stack(auxiliary_losses).mean()
|
| 510 |
+
total_loss = main_loss + effective_mtp_weight * mtp_loss
|
| 511 |
+
else:
|
| 512 |
+
total_loss = main_loss
|
| 513 |
+
if z_loss is not None:
|
| 514 |
+
total_loss = total_loss + z_loss_coefficient * z_loss
|
| 515 |
+
|
| 516 |
+
output_states = (
|
| 517 |
+
hidden_states
|
| 518 |
+
if logits_to_keep is None
|
| 519 |
+
else hidden_states[:, -logits_to_keep:, :]
|
| 520 |
+
)
|
| 521 |
+
logits = self.project_to_vocabulary(output_states) if return_logits else None
|
| 522 |
+
return TinyGDNOutput(
|
| 523 |
+
loss=total_loss,
|
| 524 |
+
logits=logits,
|
| 525 |
+
main_loss=main_loss,
|
| 526 |
+
mtp_loss=mtp_loss,
|
| 527 |
+
z_loss=z_loss,
|
| 528 |
+
hidden_states=hidden_states if return_hidden_states else None,
|
| 529 |
+
)
|
| 530 |
+
|
| 531 |
+
def parameter_report(self) -> dict[str, int]:
|
| 532 |
+
total = sum(parameter.numel() for parameter in self.parameters())
|
| 533 |
+
mtp = sum(parameter.numel() for parameter in self.mtp_adapters.parameters())
|
| 534 |
+
embeddings = self.embed_tokens.weight.numel()
|
| 535 |
+
return {
|
| 536 |
+
"deployable_core": total - mtp,
|
| 537 |
+
"training_total": total,
|
| 538 |
+
"embedding": embeddings,
|
| 539 |
+
"mtp_auxiliary": mtp,
|
| 540 |
+
"non_embedding_core": total - mtp - embeddings,
|
| 541 |
+
}
|
| 542 |
+
|
| 543 |
+
def save_checkpoint(self, output_dir: Path) -> None:
|
| 544 |
+
output_dir.mkdir(parents=True, exist_ok=True)
|
| 545 |
+
self.config.save_json(output_dir / "config.json")
|
| 546 |
+
save_model(self, output_dir / "model.safetensors")
|
| 547 |
+
|
| 548 |
+
@classmethod
|
| 549 |
+
def from_checkpoint(
|
| 550 |
+
cls,
|
| 551 |
+
checkpoint_dir: Path,
|
| 552 |
+
*,
|
| 553 |
+
device: str | torch.device = "cpu",
|
| 554 |
+
dtype: torch.dtype | None = None,
|
| 555 |
+
) -> TinyGDNForCausalLM:
|
| 556 |
+
config = TinyGDNConfig.from_json(checkpoint_dir / "config.json")
|
| 557 |
+
model = cls(config).to(device=device, dtype=dtype)
|
| 558 |
+
load_model(model, checkpoint_dir / "model.safetensors", device=str(device))
|
| 559 |
+
return model
|
| 560 |
+
|
| 561 |
+
def extra_repr(self) -> str:
|
| 562 |
+
report = self.parameter_report()
|
| 563 |
+
return (
|
| 564 |
+
f"core_parameters={report['deployable_core']:,}, "
|
| 565 |
+
f"training_parameters={report['training_total']:,}"
|
| 566 |
+
)
|
| 567 |
+
|
| 568 |
+
def get_architecture_metadata(self) -> dict[str, Any]:
|
| 569 |
+
return {
|
| 570 |
+
"architecture": self.config.architecture,
|
| 571 |
+
"layer_types": list(self.config.layer_types),
|
| 572 |
+
"effective_num_layers": self.config.effective_num_layers,
|
| 573 |
+
"shared_layer_indices": list(self.config.shared_layer_indices),
|
| 574 |
+
"parameter_report": self.parameter_report(),
|
| 575 |
+
"features": [
|
| 576 |
+
"32-layer deep-thin parameter allocation",
|
| 577 |
+
"Gated DeltaNet-2 recurrent memory",
|
| 578 |
+
"3:1 recurrent-to-full-attention hybrid",
|
| 579 |
+
"gated grouped-query attention",
|
| 580 |
+
"QK normalization",
|
| 581 |
+
"partial rotary embeddings",
|
| 582 |
+
"zero-centered RMSNorm",
|
| 583 |
+
"SwiGLU",
|
| 584 |
+
"tied input-output embeddings",
|
| 585 |
+
"optional multi-token prediction auxiliaries",
|
| 586 |
+
],
|
| 587 |
+
}
|
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": "{%- for message in messages -%}\n{%- if loop.first -%}{{- bos_token -}}{%- endif -%}\n{{- '<|im_start|>' + message['role'] + '\\n' -}}\n{%- if message['role'] == 'tool' -%}{{- '<|tool_response|>\\n' -}}{%- endif -%}\n{%- if message['content'] is string -%}\n{{- message['content'] -}}\n{%- elif message['content'] is iterable -%}\n{%- for item in message['content'] -%}\n{%- if item['type'] == 'text' -%}{{- item['text'] -}}{%- endif -%}\n{%- endfor -%}\n{%- endif -%}\n{%- if message['tool_calls'] is defined and message['tool_calls'] -%}\n{{- '\\n<|tool_call|>\\n' + (message['tool_calls'] | tojson) -}}\n{%- endif -%}\n{{- '<|im_end|>\\n' -}}\n{%- endfor -%}\n{%- if add_generation_prompt -%}{{- '<|im_start|>assistant\\n' -}}{%- endif -%}",
|
| 12 |
+
"extra_special_tokens": [
|
| 13 |
+
"<|im_start|>",
|
| 14 |
+
"<|im_end|>",
|
| 15 |
+
"<|tool_call|>",
|
| 16 |
+
"<|tool_response|>",
|
| 17 |
+
"<|think_start|>",
|
| 18 |
+
"<|think_end|>",
|
| 19 |
+
"<|fim_prefix|>",
|
| 20 |
+
"<|fim_middle|>",
|
| 21 |
+
"<|fim_suffix|>",
|
| 22 |
+
"<|fim_pad|>",
|
| 23 |
+
"<|reserved_000|>",
|
| 24 |
+
"<|reserved_001|>",
|
| 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-00010619",
|
| 3 |
+
"ema": {
|
| 4 |
+
"assistant_targets": 91313,
|
| 5 |
+
"batches": 16,
|
| 6 |
+
"elapsed_seconds": 1.9623924390034517,
|
| 7 |
+
"loss": 1.5321617420986058,
|
| 8 |
+
"perplexity": 4.628170928002891,
|
| 9 |
+
"targets_per_second": 46531.46750115428
|
| 10 |
+
},
|
| 11 |
+
"normal": {
|
| 12 |
+
"assistant_targets": 91313,
|
| 13 |
+
"batches": 16,
|
| 14 |
+
"elapsed_seconds": 1.9647464910085546,
|
| 15 |
+
"loss": 1.5321617420986058,
|
| 16 |
+
"perplexity": 4.628170928002891,
|
| 17 |
+
"targets_per_second": 46475.716036589896
|
| 18 |
+
},
|
| 19 |
+
"optimizer_step": 10619,
|
| 20 |
+
"type": "sft_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 |
+
]
|