kerzgrr commited on
Commit
fba158e
·
verified ·
1 Parent(s): 4748a29

Upload Haiku-base (pretrain EMA @ step 8400)

Browse files
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
+ [![Model](https://img.shields.io/badge/Model-~655M_params-blue)](.)
30
+ [![Stage](https://img.shields.io/badge/Stage-Pretrain_(base)-orange.svg)](.)
31
+ [![License](https://img.shields.io/badge/License-Apache_2.0-green.svg)](LICENSE)
32
+ [![Architecture](https://img.shields.io/badge/Arch-KDA_+_Gated_MLA-purple.svg)](.)
33
+ [![Demo](https://img.shields.io/badge/Space-haiku--demo-indigo.svg)](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
+ ]