kerzgrr commited on
Commit
6bcec9b
·
verified ·
1 Parent(s): 7fa0c0e

Upload Monostich-2 SFT EMA (smoltalk+hermes 1 epoch, step 10619, 21.70h)

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