tturing commited on
Commit
4e9d3a8
·
verified ·
1 Parent(s): bf85774

Add Q35KDA_fp8base_ffn8_qkvo8_fp11_hd128_11785_fm128: KDA:GQA 1.7B, SD (stale) arm, seed 11785 (standalone, trust_remote_code)

Browse files
README.md ADDED
@@ -0,0 +1,88 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ tags:
3
+ - kimi-delta-attention
4
+ - kda
5
+ - hybrid
6
+ - fp8
7
+ - deltamatching
8
+ ---
9
+
10
+ # DM-KDAGQA-1.7B-SD
11
+
12
+ A 1.66 B-parameter KDA : GQA hybrid (Kimi Delta Attention recurrent layers, grouped-query attention) pretrained from
13
+ scratch on 30 B tokens with naive FP8 attention: FlashMatch FP8 attention with the stale delta (`sb_mode=fp8base`). It is the **SD** (stale) arm of a study in which three models were trained
14
+ identically except for the attention's training precision:
15
+
16
+ | repo | arm | attention in training | FP8 GEMMs |
17
+ |---|---|---|---|
18
+ | [DM-KDAGQA-1.7B-MP](https://huggingface.co/tturing/DM-KDAGQA-1.7B-MP) | clean | bf16 cuDNN SDPA (standard BF16/FP32 mixed precision) | FFN |
19
+ | [DM-KDAGQA-1.7B-SD](https://huggingface.co/tturing/DM-KDAGQA-1.7B-SD) | stale | FlashMatch FP8 attention, stale delta (naive FP8) | FFN + attention q/k/v/o |
20
+ | [DM-KDAGQA-1.7B-DM](https://huggingface.co/tturing/DM-KDAGQA-1.7B-DM) | match | FlashMatch FP8 attention, DeltaMatching (matched delta) | FFN + attention q/k/v/o |
21
+
22
+ ## Model
23
+
24
+ | | |
25
+ |---|---|
26
+ | layers | 24 = [KDA, KDA, KDA, GQA] × 6 |
27
+ | width | d_model 2048, SwiGLU FFN 8,064, tied embeddings |
28
+ | KDA | Kimi Delta Attention (flash-linear-attention, chunk mode), 16 heads × 128, value width = key width, short convolution, per-channel decay gate, gated output |
29
+ | attention | GQA, 16 query / 4 KV heads × 128, qk-norm, partial RoPE 0.25 (θ = 1e7), gated output |
30
+ | vocabulary, context | Llama-2 32k tokenizer, 8,192 tokens |
31
+ | parameters | 1,664,776,224 (bf16 safetensors) |
32
+
33
+ The export runs attention through bf16 SDPA in every arm, so the three repos share one architecture and differ only in
34
+ their weights. `config.train.json` records the training-time layer config (the FP8 arms train their attention layers
35
+ through the `sagebwd` FlashMatch mixer). The KDA layers are bf16 in all three arms.
36
+
37
+ ## Training
38
+
39
+ | | |
40
+ |---|---|
41
+ | data | Nemotron-CC (`nemotron_cc_v2d1_hq_dqa`), packed 8,192-token sequences |
42
+ | budget | 28,610 steps × global batch 128 × 8,192 tokens = 30.0 B tokens |
43
+ | optimizer | AdamW, β (0.9, 0.95), ε 1e-8, weight decay 0.1, gradient clip 1.0, z-loss 1e-4 |
44
+ | schedule | WSD: 3.33 % warmup to 2.4e-3, constant, linear decay over the last 20 % |
45
+ | seed | 11785 (initialization only; every run in the study reads the same data in the same order) |
46
+ | hardware | 8 × H200 |
47
+
48
+ ## Evaluation (seed 11785)
49
+
50
+ | model | val CE | RULER | CSense-9 | Extract | TriviaQA | MMLU | MQAR |
51
+ |---|---|---|---|---|---|---|---|
52
+ | MP (clean) | 1.3990 | 58.66 | 0.5994 | 0.6829 | 0.1764 | 0.3472 | 0.1242 |
53
+ | **SD (stale)** | **1.7124** | **34.58** | **0.5072** | **0.5546** | **0.0444** | **0.2646** | **0.1337** |
54
+ | DM (match) | 1.3989 | 55.61 | 0.5993 | 0.6556 | 0.1851 | 0.3374 | 0.1369 |
55
+
56
+ val CE: nats/token on 2,000 held-out 8,192-token windows (lower is better). RULER: 13 tasks at 4k and 8k, mean of the
57
+ two lengths. CSense-9: LAMBADA, HellaSwag, PIQA, ARC-e, ARC-c, SciQ, OpenBookQA, WinoGrande, COPA. Extract: SWDE, FDA,
58
+ SQuAD completion. TriviaQA: 5-shot exact match. MQAR: synthetic multi-query key-value recall. Each arm was trained
59
+ with two seeds. In this cell the stale arm drifted away from clean during training (two-seed means: val CE +0.42,
60
+ RULER 24.5 against 57.8), while the match arm equals clean on val CE (−0.0006) and stays within 0.0011 of it in train
61
+ CE from step 5,000 on; its RULER −2.1 comes from one subtask (`niah_multikey_2`). Per-seed and per-task results and the
62
+ full configuration: `report.md` in [tturing/n8t-train-curve](https://huggingface.co/datasets/tturing/n8t-train-curve).
63
+
64
+ ## Usage
65
+
66
+ The model code ships with this repo (`modeling_bqalm.py`, `configuration_bqalm.py`), so no other codebase is needed.
67
+ It needs a CUDA GPU with `torch`, `transformers` >= 5 and `flash-linear-attention` (tested with torch 2.12,
68
+ transformers 5.9.0, flash-linear-attention 0.5.0).
69
+
70
+ ```python
71
+ import torch
72
+ from transformers import AutoModelForCausalLM, AutoTokenizer
73
+
74
+ repo = "tturing/DM-KDAGQA-1.7B-SD"
75
+ model = AutoModelForCausalLM.from_pretrained(repo, dtype=torch.bfloat16, trust_remote_code=True).cuda()
76
+ tokenizer = AutoTokenizer.from_pretrained(repo, trust_remote_code=True)
77
+
78
+ inputs = tokenizer("The capital of France is", return_tensors="pt").to("cuda")
79
+ print(tokenizer.decode(model.generate(**inputs, max_new_tokens=32)[0], skip_special_tokens=True))
80
+ ```
81
+
82
+ - **Inference.** `generate()` keeps a cache (the GQA layers' K/V and the KDA recurrent and conv states), so each new
83
+ token costs one step. Greedy decoding and sampling are supported (`num_beams=1`); prompts batched together must
84
+ share one length, since the KDA layers have no padding mask.
85
+ - **Fine-tuning.** `model(input_ids=ids, labels=ids).loss` is the next-token cross-entropy, so the model trains with
86
+ the `transformers` `Trainer` in bf16. The FP8 FlashMatch attention the SD and DM arms were trained with needs a
87
+ compiled CUDA kernel that is not shipped; the exported weights are bf16 and run bf16 SDPA attention.
88
+ - **Fidelity.** The shipped code is the study's evaluation path: its logits match the original loader bit for bit.
config.json ADDED
@@ -0,0 +1,72 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "attn_impl": "sdpa",
3
+ "attn_output_gate": true,
4
+ "bqa_local_precision": "fp8",
5
+ "bqa_nvfp4_block": 16,
6
+ "bqa_remote_k_precision": "fp8",
7
+ "bqa_remote_topk": -1,
8
+ "bqa_remote_v_precision": "nvfp4",
9
+ "bqa_rotate": true,
10
+ "bqa_window": 512,
11
+ "d_model": 2048,
12
+ "ffn_mult": 2.6666666666666665,
13
+ "ffn_multiple_of": 256,
14
+ "gdn_head_dim": 128,
15
+ "gdn_num_heads": 16,
16
+ "gdn_num_v_heads": null,
17
+ "head_dim": 128,
18
+ "hidden_size": 2048,
19
+ "intermediate_size": 8064,
20
+ "kda_allow_neg_eigval": false,
21
+ "kda_conv_size": 4,
22
+ "kda_expand_v": 1.0,
23
+ "kda_head_dim": null,
24
+ "kda_lower_bound": null,
25
+ "kda_num_heads": null,
26
+ "kda_num_v_heads": null,
27
+ "kda_safe_gate": false,
28
+ "kv_lora_rank": null,
29
+ "layer_mixers": [
30
+ "kda",
31
+ "kda",
32
+ "kda",
33
+ "gqa"
34
+ ],
35
+ "mamba2_chunk_size": 256,
36
+ "mamba2_d_state": 128,
37
+ "mamba2_expand": 2,
38
+ "mamba2_headdim": null,
39
+ "mamba2_ngroups": 1,
40
+ "max_position_embeddings": 8192,
41
+ "max_seq_len": 8192,
42
+ "mixer": "gqa",
43
+ "model_type": "bqalm",
44
+ "n_heads": 16,
45
+ "n_kv_heads": 4,
46
+ "n_layers": 24,
47
+ "nope": false,
48
+ "norm_eps": 1e-06,
49
+ "num_attention_heads": 16,
50
+ "num_hidden_layers": 24,
51
+ "partial_rotary_factor": 0.25,
52
+ "qk_norm": true,
53
+ "rms_norm_in_fp32": true,
54
+ "rope_parameters": {
55
+ "partial_rotary_factor": 0.25,
56
+ "rope_theta": 10000000,
57
+ "rope_type": "default"
58
+ },
59
+ "rope_theta": 10000000,
60
+ "tie_embeddings": true,
61
+ "tie_word_embeddings": true,
62
+ "transformers_version": "5.9.0",
63
+ "vocab_size": 32000,
64
+ "z_loss_weight": 0.0,
65
+ "architectures": [
66
+ "BqaLMForCausalLM"
67
+ ],
68
+ "auto_map": {
69
+ "AutoConfig": "configuration_bqalm.BqaLMConfig",
70
+ "AutoModelForCausalLM": "modeling_bqalm.BqaLMForCausalLM"
71
+ }
72
+ }
config.train.json ADDED
@@ -0,0 +1,66 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "attn_impl": "sdpa",
3
+ "attn_output_gate": true,
4
+ "bqa_local_precision": "fp8",
5
+ "bqa_nvfp4_block": 16,
6
+ "bqa_remote_k_precision": "fp8",
7
+ "bqa_remote_topk": -1,
8
+ "bqa_remote_v_precision": "nvfp4",
9
+ "bqa_rotate": true,
10
+ "bqa_window": 512,
11
+ "d_model": 2048,
12
+ "ffn_mult": 2.6666666666666665,
13
+ "ffn_multiple_of": 256,
14
+ "gdn_head_dim": 128,
15
+ "gdn_num_heads": 16,
16
+ "gdn_num_v_heads": null,
17
+ "head_dim": 128,
18
+ "hidden_size": 2048,
19
+ "intermediate_size": 8064,
20
+ "kda_allow_neg_eigval": false,
21
+ "kda_conv_size": 4,
22
+ "kda_expand_v": 1.0,
23
+ "kda_head_dim": null,
24
+ "kda_lower_bound": null,
25
+ "kda_num_heads": null,
26
+ "kda_num_v_heads": null,
27
+ "kda_safe_gate": false,
28
+ "kv_lora_rank": null,
29
+ "layer_mixers": [
30
+ "kda",
31
+ "kda",
32
+ "kda",
33
+ "sagebwd"
34
+ ],
35
+ "mamba2_chunk_size": 256,
36
+ "mamba2_d_state": 128,
37
+ "mamba2_expand": 2,
38
+ "mamba2_headdim": null,
39
+ "mamba2_ngroups": 1,
40
+ "max_position_embeddings": 8192,
41
+ "max_seq_len": 8192,
42
+ "mixer": "gqa",
43
+ "model_type": "bqalm",
44
+ "n_heads": 16,
45
+ "n_kv_heads": 4,
46
+ "n_layers": 24,
47
+ "nope": false,
48
+ "norm_eps": 1e-06,
49
+ "num_attention_heads": 16,
50
+ "num_hidden_layers": 24,
51
+ "partial_rotary_factor": 0.25,
52
+ "qk_norm": true,
53
+ "rms_norm_in_fp32": true,
54
+ "rope_parameters": {
55
+ "partial_rotary_factor": 0.25,
56
+ "rope_theta": 10000000,
57
+ "rope_type": "default"
58
+ },
59
+ "rope_theta": 10000000,
60
+ "sb_impl": "fa3fp11",
61
+ "tie_embeddings": true,
62
+ "tie_word_embeddings": true,
63
+ "transformers_version": "5.9.0",
64
+ "vocab_size": 32000,
65
+ "z_loss_weight": 0.0
66
+ }
configuration_bqalm.py ADDED
@@ -0,0 +1,272 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """HF `PretrainedConfig` for BqaLM, the hybrid decoder of the DeltaMatching study.
2
+
3
+ Self-contained copy of the bqa codebase's `src/pretrain/hf/configuration_bqalm.py`, with the backbone's
4
+ `ModelConfig` (`src/pretrain/modeling/config.py`) vendored below so the checkpoint loads through
5
+ `trust_remote_code` without the bqa source tree. Field names, defaults and `to_model_config` are unchanged, so a
6
+ `config.json` written by the bqa exporter round-trips exactly.
7
+ """
8
+ from __future__ import annotations
9
+
10
+ from dataclasses import dataclass, field
11
+
12
+ from transformers import PretrainedConfig
13
+
14
+
15
+ @dataclass
16
+ class ModelConfig:
17
+ """Backbone architecture (the bqa `ModelConfig`). Only the fields the shipped mixers read matter here; the rest
18
+ are carried so the exporter's configs keep their exact meaning."""
19
+ # core dims
20
+ vocab_size: int = 32768
21
+ d_model: int = 1024
22
+ n_layers: int = 24
23
+ n_heads: int = 16
24
+ n_kv_heads: int = 4 # GQA: n_heads % n_kv_heads == 0; == n_heads is MHA
25
+ head_dim: int | None = None # default d_model // n_heads
26
+ # FFN (SwiGLU)
27
+ intermediate_size: int | None = None
28
+ ffn_mult: float = 8.0 / 3.0
29
+ ffn_multiple_of: int = 256
30
+ # positional / norm
31
+ rope_theta: float = 10000.0
32
+ rope_scaling: dict | None = None # {"type": "yarn", ...} for YaRN context extension, else None
33
+ partial_rotary_factor: float = 1.0 # rotate the leading head_dim * factor channels only
34
+ nope: bool = False # identity rotation (no positional encoding in attention)
35
+ norm_eps: float = 1e-5
36
+ max_seq_len: int = 2048
37
+ # mixers
38
+ mixer: str = "gqa"
39
+ attn_impl: str = "auto"
40
+ qk_norm: bool = False # per-head RMSNorm on q, k before RoPE
41
+ attn_output_gate: bool = False # q_proj emits 2 * q_dim; out * sigmoid(gate) before o_proj
42
+ layer_mixers: list[str] | None = None # per-layer pattern, repeated over n_layers
43
+ gdn_head_dim: int | None = None
44
+ gdn_num_heads: int | None = None
45
+ gdn_num_v_heads: int | None = None
46
+ kv_lora_rank: int | None = None
47
+ mamba2_headdim: int | None = None
48
+ mamba2_d_state: int = 128
49
+ mamba2_expand: int = 2
50
+ mamba2_ngroups: int = 1
51
+ mamba2_chunk_size: int = 256
52
+ kda_head_dim: int | None = None
53
+ kda_num_heads: int | None = None
54
+ kda_num_v_heads: int | None = None
55
+ kda_expand_v: float = 1.0
56
+ kda_conv_size: int = 4
57
+ kda_allow_neg_eigval: bool = False
58
+ kda_safe_gate: bool = False
59
+ kda_lower_bound: float | None = None
60
+ # numerics
61
+ tie_embeddings: bool = True
62
+ attn_dropout: float = 0.0
63
+ resid_dropout: float = 0.0
64
+ initializer_range: float = 0.02
65
+ z_loss_weight: float = 1e-4
66
+ rms_norm_in_fp32: bool = True
67
+ fused_rmsnorm: bool = False
68
+ fused_rope: bool = False
69
+ fused_swiglu: bool = False
70
+ # training-kernel and BQA-mixer knobs, carried for config fidelity only (not read by the shipped mixers)
71
+ sb_mode: str = "uniform"
72
+ sb_window: int = 512
73
+ sb_sink: bool = True
74
+ sb_impl: str = "triton"
75
+ sb_t0_dedup: bool = False
76
+ sb_fused_producers: bool = False
77
+ sb_fuse_tier2: str = "off"
78
+ hs_protect: str = ""
79
+ hs_impl: str = "compose"
80
+ bqa_window: int = 512
81
+ bqa_rotate: bool = True
82
+ bqa_nvfp4_block: int = 16
83
+ bqa_remote_topk: int = -1
84
+ bqa_local_precision: str = "fp8"
85
+ bqa_remote_k_precision: str = "fp8"
86
+ bqa_remote_v_precision: str = "nvfp4"
87
+ _resolved: bool = field(default=False, repr=False)
88
+
89
+ def __post_init__(self):
90
+ if self.head_dim is None:
91
+ assert self.d_model % self.n_heads == 0, "d_model must divide by n_heads when head_dim is None"
92
+ self.head_dim = self.d_model // self.n_heads
93
+ assert self.n_heads % self.n_kv_heads == 0, (
94
+ f"n_heads ({self.n_heads}) must be divisible by n_kv_heads ({self.n_kv_heads})")
95
+ if self.intermediate_size is None:
96
+ raw = self.ffn_mult * self.d_model
97
+ m = self.ffn_multiple_of
98
+ self.intermediate_size = int(((int(raw) + m - 1) // m) * m)
99
+ assert 0.0 < self.partial_rotary_factor <= 1.0, (
100
+ f"partial_rotary_factor must be in (0, 1], got {self.partial_rotary_factor}")
101
+ assert self.rotary_dim % 2 == 0 and self.rotary_dim > 0, (
102
+ f"rotary_dim = head_dim({self.head_dim}) * partial_rotary_factor"
103
+ f"({self.partial_rotary_factor}) = {self.rotary_dim}, which must be a positive even number")
104
+ if self.rotary_dim != self.head_dim and self.fused_rope:
105
+ self.fused_rope = False
106
+ self._resolved = True
107
+
108
+ @property
109
+ def rotary_dim(self) -> int:
110
+ return int(self.head_dim * self.partial_rotary_factor)
111
+
112
+ @property
113
+ def n_rep(self) -> int:
114
+ return self.n_heads // self.n_kv_heads
115
+
116
+ @property
117
+ def q_dim(self) -> int:
118
+ return self.n_heads * self.head_dim
119
+
120
+ @property
121
+ def kv_dim(self) -> int:
122
+ return self.n_kv_heads * self.head_dim
123
+
124
+ def mixer_for_layer(self, layer_idx: int) -> str:
125
+ if self.layer_mixers is not None:
126
+ return self.layer_mixers[layer_idx % len(self.layer_mixers)]
127
+ return self.mixer
128
+
129
+
130
+ class BqaLMConfig(PretrainedConfig):
131
+ model_type = "bqalm"
132
+
133
+ def __init__(
134
+ self,
135
+ vocab_size: int = 50257,
136
+ d_model: int = 1024,
137
+ n_layers: int = 24,
138
+ n_heads: int = 16,
139
+ n_kv_heads: int = 4,
140
+ head_dim: int | None = None,
141
+ intermediate_size: int | None = None,
142
+ ffn_mult: float = 8.0 / 3.0,
143
+ ffn_multiple_of: int = 256,
144
+ rope_theta: float = 10000.0,
145
+ rope_scaling: dict | None = None,
146
+ norm_eps: float = 1e-5,
147
+ rms_norm_in_fp32: bool = True,
148
+ qk_norm: bool = False,
149
+ attn_output_gate: bool = False,
150
+ partial_rotary_factor: float = 1.0,
151
+ nope: bool = False,
152
+ gdn_head_dim: int | None = None,
153
+ gdn_num_heads: int | None = None,
154
+ gdn_num_v_heads: int | None = None,
155
+ kv_lora_rank: int | None = None,
156
+ mamba2_headdim: int | None = None,
157
+ mamba2_d_state: int = 128,
158
+ mamba2_expand: int = 2,
159
+ mamba2_ngroups: int = 1,
160
+ mamba2_chunk_size: int = 256,
161
+ kda_head_dim: int | None = None,
162
+ kda_num_heads: int | None = None,
163
+ kda_num_v_heads: int | None = None,
164
+ kda_expand_v: float = 1.0,
165
+ kda_conv_size: int = 4,
166
+ kda_allow_neg_eigval: bool = False,
167
+ kda_safe_gate: bool = False,
168
+ kda_lower_bound: float | None = None,
169
+ max_seq_len: int = 2048,
170
+ mixer: str = "gqa",
171
+ attn_impl: str = "auto",
172
+ layer_mixers: list[str] | None = None,
173
+ tie_embeddings: bool = True,
174
+ z_loss_weight: float = 0.0,
175
+ bqa_window: int = 512,
176
+ bqa_remote_topk: int = 512,
177
+ bqa_local_precision: str = "fp8",
178
+ bqa_remote_k_precision: str = "fp8",
179
+ bqa_remote_v_precision: str = "nvfp4",
180
+ bqa_rotate: bool = True,
181
+ bqa_nvfp4_block: int = 16,
182
+ sb_impl: str = "triton",
183
+ **kwargs,
184
+ ):
185
+ self.vocab_size = vocab_size
186
+ self.d_model = d_model
187
+ self.n_layers = n_layers
188
+ self.n_heads = n_heads
189
+ self.n_kv_heads = n_kv_heads
190
+ self.head_dim = head_dim
191
+ self.intermediate_size = intermediate_size
192
+ self.ffn_mult = ffn_mult
193
+ self.ffn_multiple_of = ffn_multiple_of
194
+ self.rope_theta = rope_theta
195
+ self.rope_scaling = rope_scaling
196
+ self.norm_eps = norm_eps
197
+ self.rms_norm_in_fp32 = rms_norm_in_fp32
198
+ self.qk_norm = qk_norm
199
+ self.attn_output_gate = attn_output_gate
200
+ self.partial_rotary_factor = partial_rotary_factor
201
+ self.nope = bool(nope)
202
+ self.gdn_head_dim = gdn_head_dim
203
+ self.gdn_num_heads = gdn_num_heads
204
+ self.gdn_num_v_heads = gdn_num_v_heads
205
+ self.kv_lora_rank = kv_lora_rank
206
+ self.mamba2_headdim = mamba2_headdim
207
+ self.mamba2_d_state = mamba2_d_state
208
+ self.mamba2_expand = mamba2_expand
209
+ self.mamba2_ngroups = mamba2_ngroups
210
+ self.mamba2_chunk_size = mamba2_chunk_size
211
+ self.kda_head_dim = kda_head_dim
212
+ self.kda_num_heads = kda_num_heads
213
+ self.kda_num_v_heads = kda_num_v_heads
214
+ self.kda_expand_v = kda_expand_v
215
+ self.kda_conv_size = kda_conv_size
216
+ self.kda_allow_neg_eigval = kda_allow_neg_eigval
217
+ self.kda_safe_gate = kda_safe_gate
218
+ self.kda_lower_bound = kda_lower_bound
219
+ self.max_seq_len = max_seq_len
220
+ # set before super().__init__: transformers' rope validation reads max_position_embeddings during init
221
+ self.max_position_embeddings = max_seq_len
222
+ self.mixer = mixer
223
+ self.attn_impl = attn_impl
224
+ self.layer_mixers = layer_mixers
225
+ self.tie_embeddings = tie_embeddings
226
+ self.z_loss_weight = z_loss_weight
227
+ self.bqa_window = bqa_window
228
+ self.bqa_remote_topk = bqa_remote_topk
229
+ self.bqa_local_precision = bqa_local_precision
230
+ self.bqa_remote_k_precision = bqa_remote_k_precision
231
+ self.bqa_remote_v_precision = bqa_remote_v_precision
232
+ self.bqa_rotate = bqa_rotate
233
+ self.bqa_nvfp4_block = bqa_nvfp4_block
234
+ self.sb_impl = sb_impl
235
+ kwargs.setdefault("max_position_embeddings", max_seq_len)
236
+ kwargs.setdefault("hidden_size", d_model)
237
+ kwargs.setdefault("num_hidden_layers", n_layers)
238
+ kwargs.setdefault("num_attention_heads", n_heads)
239
+ kwargs.setdefault("tie_word_embeddings", tie_embeddings)
240
+ super().__init__(**kwargs)
241
+
242
+ def to_model_config(self) -> ModelConfig:
243
+ return ModelConfig(
244
+ vocab_size=self.vocab_size, d_model=self.d_model, n_layers=self.n_layers,
245
+ n_heads=self.n_heads, n_kv_heads=self.n_kv_heads, head_dim=self.head_dim,
246
+ intermediate_size=self.intermediate_size, ffn_mult=self.ffn_mult,
247
+ ffn_multiple_of=self.ffn_multiple_of, rope_theta=self.rope_theta,
248
+ rope_scaling=self.rope_scaling,
249
+ norm_eps=self.norm_eps, rms_norm_in_fp32=self.rms_norm_in_fp32, qk_norm=self.qk_norm,
250
+ attn_output_gate=self.attn_output_gate,
251
+ partial_rotary_factor=self.partial_rotary_factor,
252
+ nope=bool(getattr(self, "nope", False)),
253
+ gdn_head_dim=self.gdn_head_dim, gdn_num_heads=self.gdn_num_heads,
254
+ gdn_num_v_heads=self.gdn_num_v_heads,
255
+ kv_lora_rank=getattr(self, "kv_lora_rank", None),
256
+ mamba2_headdim=getattr(self, "mamba2_headdim", None), mamba2_d_state=getattr(self, "mamba2_d_state", 128),
257
+ mamba2_expand=getattr(self, "mamba2_expand", 2), mamba2_ngroups=getattr(self, "mamba2_ngroups", 1),
258
+ mamba2_chunk_size=getattr(self, "mamba2_chunk_size", 256),
259
+ kda_head_dim=getattr(self, "kda_head_dim", None), kda_num_heads=getattr(self, "kda_num_heads", None),
260
+ kda_num_v_heads=getattr(self, "kda_num_v_heads", None), kda_expand_v=getattr(self, "kda_expand_v", 1.0),
261
+ kda_conv_size=getattr(self, "kda_conv_size", 4), kda_allow_neg_eigval=getattr(self, "kda_allow_neg_eigval", False),
262
+ kda_safe_gate=getattr(self, "kda_safe_gate", False), kda_lower_bound=getattr(self, "kda_lower_bound", None),
263
+ max_seq_len=self.max_seq_len, mixer=self.mixer,
264
+ attn_impl=self.attn_impl, layer_mixers=self.layer_mixers,
265
+ tie_embeddings=self.tie_embeddings, z_loss_weight=self.z_loss_weight,
266
+ bqa_window=self.bqa_window, bqa_remote_topk=self.bqa_remote_topk,
267
+ bqa_local_precision=self.bqa_local_precision,
268
+ bqa_remote_k_precision=self.bqa_remote_k_precision,
269
+ bqa_remote_v_precision=self.bqa_remote_v_precision,
270
+ bqa_rotate=self.bqa_rotate, bqa_nvfp4_block=self.bqa_nvfp4_block,
271
+ sb_impl=self.sb_impl,
272
+ )
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:99a6c1cf7fe222756980194353a6847846ba41110cc8d70248ec8d24b41f2f7f
3
+ size 3329603472
modeling_bqalm.py ADDED
@@ -0,0 +1,630 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """BqaLM for `transformers` (remote code): the hybrid decoder of the DeltaMatching study.
2
+
3
+ Self-contained copy of the bqa codebase's inference path (`src/pretrain/modeling` + `src/pretrain/hf`) for the mixers the
4
+ published checkpoints use: GQA and MLA attention, Mamba2, GatedDeltaNet, KDA. Training-only paths (FP8 FlashMatch
5
+ attention, fused FP8 producers, TransformerEngine, the Liger fused kernels) are left out; what remains is the eager /
6
+ SDPA code the study's evaluation ran, so the logits match the bqa loader bit for bit. Parameter names are the
7
+ checkpoint's: `model.embed_tokens`, `model.layers.N.{input_layernorm,mixer,post_attention_layernorm,mlp}`,
8
+ `model.norm`, `lm_head` (tied to the embedding).
9
+
10
+ Requirements: torch and transformers; Mamba2 layers also need `mamba_ssm` + `causal_conv1d`, GatedDeltaNet and KDA
11
+ layers `flash-linear-attention` (`fla`). Each is imported only by the layers that use it, with an install hint if missing.
12
+
13
+ `generate()` keeps a cache (`BqaLMCache`: K/V of the attention layers, mamba_ssm's conv / SSM states, fla's
14
+ recurrent and conv states), so each new token costs one step. It supports greedy decoding and sampling (num_beams=1);
15
+ a batch of prompts must share a length, since the recurrent layers have no padding mask. An all-ones attention mask
16
+ is the same as none. A plain forward (no `use_cache`) is the uncached evaluation path.
17
+ """
18
+ from __future__ import annotations
19
+
20
+ import math
21
+ from dataclasses import dataclass
22
+
23
+ import torch
24
+ import torch.nn as nn
25
+ import torch.nn.functional as F
26
+ from transformers import GenerationMixin, PreTrainedModel
27
+ from transformers.utils import ModelOutput
28
+
29
+ from .configuration_bqalm import BqaLMConfig, ModelConfig
30
+
31
+
32
+ # ----------------------------------------------------------------------------------------------- primitive layers
33
+ class RMSNorm(nn.Module):
34
+ """x / sqrt(mean(x^2) + eps) * w, reduced in fp32 when `in_fp32` (then cast back before the scale)."""
35
+
36
+ def __init__(self, dim: int, eps: float = 1e-5, in_fp32: bool = True):
37
+ super().__init__()
38
+ self.eps = eps
39
+ self.in_fp32 = in_fp32
40
+ self.weight = nn.Parameter(torch.ones(dim))
41
+
42
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
43
+ dtype = x.dtype
44
+ if self.in_fp32:
45
+ x = x.float()
46
+ var = x.pow(2).mean(dim=-1, keepdim=True)
47
+ x = x * torch.rsqrt(var + self.eps)
48
+ return (self.weight * x.to(dtype)) if self.in_fp32 else (self.weight * x)
49
+
50
+
51
+ def _is_yarn(rope_scaling: dict | None) -> bool:
52
+ return bool(rope_scaling) and str(rope_scaling.get("type", rope_scaling.get("rope_type", ""))).lower() == "yarn"
53
+
54
+
55
+ def yarn_mscale(rope_scaling: dict | None) -> float:
56
+ """YaRN attention temperature applied to post-RoPE q (1.0 == no scaling)."""
57
+ if not _is_yarn(rope_scaling):
58
+ return 1.0
59
+ if rope_scaling.get("mscale") is not None:
60
+ return float(rope_scaling["mscale"])
61
+ factor = float(rope_scaling["factor"])
62
+ if factor <= 1.0:
63
+ return 1.0
64
+ return 0.1 * math.log(factor) + 1.0
65
+
66
+
67
+ def _yarn_find_dim(num_rotations: float, head_dim: int, theta: float, max_pos: int) -> float:
68
+ return (head_dim * math.log(max_pos / (num_rotations * 2 * math.pi))) / (2 * math.log(theta))
69
+
70
+
71
+ def yarn_inv_freq(head_dim: int, theta: float, rope_scaling: dict, device, dtype=torch.float32) -> torch.Tensor:
72
+ """NTK-by-parts interpolated inv_freq, shape [head_dim / 2] (fp32)."""
73
+ factor = float(rope_scaling["factor"])
74
+ orig_max = int(rope_scaling.get("original_max_position_embeddings", 8192))
75
+ beta_fast = float(rope_scaling.get("beta_fast", 32))
76
+ beta_slow = float(rope_scaling.get("beta_slow", 1))
77
+ pos_freqs = theta ** (torch.arange(0, head_dim, 2, device=device, dtype=torch.float32) / head_dim)
78
+ inv_freq_extrap = 1.0 / pos_freqs
79
+ inv_freq_interp = 1.0 / (factor * pos_freqs)
80
+ low = math.floor(_yarn_find_dim(beta_fast, head_dim, theta, orig_max))
81
+ high = math.ceil(_yarn_find_dim(beta_slow, head_dim, theta, orig_max))
82
+ low = max(low, 0)
83
+ high = min(high, head_dim // 2 - 1)
84
+ if low == high:
85
+ high += 0.001
86
+ ramp = (torch.arange(head_dim // 2, device=device, dtype=torch.float32) - low) / (high - low)
87
+ ramp = torch.clamp(ramp, 0.0, 1.0)
88
+ extrap_factor = 1.0 - ramp
89
+ inv = inv_freq_interp * (1.0 - extrap_factor) + inv_freq_extrap * extrap_factor
90
+ return inv.to(dtype)
91
+
92
+
93
+ class RotaryEmbedding(nn.Module):
94
+ """cos / sin computed per forward from inv_freq (NeoX layout, fp32), over the rotated width `rotary_dim`.
95
+ No inv_freq buffer on purpose: `from_pretrained` would materialize a non-persistent buffer uninitialized."""
96
+
97
+ def __init__(self, head_dim: int, max_seq_len: int, theta: float = 10000.0,
98
+ rope_scaling: dict | None = None, rotary_dim: int | None = None):
99
+ super().__init__()
100
+ head_dim = int(rotary_dim) if rotary_dim else head_dim
101
+ assert head_dim % 2 == 0, "RoPE needs an even head_dim"
102
+ self.head_dim = head_dim
103
+ self.theta = theta
104
+ self.max_seq_len = max_seq_len
105
+ self.rope_scaling = rope_scaling
106
+ self._use_yarn = _is_yarn(rope_scaling)
107
+
108
+ def forward(self, position_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
109
+ if self._use_yarn:
110
+ inv = yarn_inv_freq(self.head_dim, self.theta, self.rope_scaling, position_ids.device)
111
+ else:
112
+ inv = 1.0 / (self.theta ** (torch.arange(0, self.head_dim, 2, device=position_ids.device,
113
+ dtype=torch.float32) / self.head_dim))
114
+ freqs = torch.einsum("bt,d->btd", position_ids.float(), inv)
115
+ emb = torch.cat([freqs, freqs], dim=-1)
116
+ return emb.cos(), emb.sin()
117
+
118
+
119
+ def rotate_half(x: torch.Tensor) -> torch.Tensor:
120
+ x1, x2 = x.chunk(2, dim=-1)
121
+ return torch.cat([-x2, x1], dim=-1)
122
+
123
+
124
+ def apply_rotary(q: torch.Tensor, k: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor):
125
+ """q, k: [B, H, T, D]; cos, sin: [B, T, R] with R <= D (partial RoPE rotates the leading R channels)."""
126
+ cos = cos.unsqueeze(1)
127
+ sin = sin.unsqueeze(1)
128
+ rd = cos.shape[-1]
129
+ if rd == q.shape[-1]:
130
+ qf, kf = q.float(), k.float()
131
+ q_out = qf * cos + rotate_half(qf) * sin
132
+ k_out = kf * cos + rotate_half(kf) * sin
133
+ return q_out.to(q.dtype), k_out.to(k.dtype)
134
+
135
+ def _split_rot(x):
136
+ xr, xp = x[..., :rd].float(), x[..., rd:]
137
+ out = xr * cos + rotate_half(xr) * sin
138
+ return torch.cat([out.to(x.dtype), xp], dim=-1)
139
+ return _split_rot(q), _split_rot(k)
140
+
141
+
142
+ def q_proj_out_features(cfg: ModelConfig) -> int:
143
+ """q_proj width, doubled when the attention output gate is on (Qwen3.5 / Qwen3-Next idiom)."""
144
+ return cfg.q_dim * (2 if cfg.attn_output_gate else 1)
145
+
146
+
147
+ def split_q_gate(qg: torch.Tensor, n_heads: int, head_dim: int, gated: bool):
148
+ """[B, T, q_dim * (2 if gated)] -> (q, gate), each [B, T, n_heads, head_dim]; the split is per head."""
149
+ if not gated:
150
+ return qg.view(*qg.shape[:2], n_heads, head_dim), None
151
+ q, g = qg.view(*qg.shape[:2], n_heads, 2 * head_dim).chunk(2, dim=-1)
152
+ return q, g
153
+
154
+
155
+ def apply_output_gate(o: torch.Tensor, gate: torch.Tensor | None) -> torch.Tensor:
156
+ if gate is None:
157
+ return o
158
+ return o * torch.sigmoid(gate.reshape(o.shape).to(o.dtype))
159
+
160
+
161
+ class SwiGLUMLP(nn.Module):
162
+ """down(silu(gate(x)) * up(x))."""
163
+
164
+ def __init__(self, d_model: int, intermediate_size: int, dropout: float = 0.0):
165
+ super().__init__()
166
+ self.gate_proj = nn.Linear(d_model, intermediate_size, bias=False)
167
+ self.up_proj = nn.Linear(d_model, intermediate_size, bias=False)
168
+ self.down_proj = nn.Linear(intermediate_size, d_model, bias=False)
169
+ self.drop = nn.Dropout(dropout) if dropout > 0 else nn.Identity()
170
+
171
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
172
+ return self.drop(self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)))
173
+
174
+
175
+ # ------------------------------------------------------------------------------------------------ generation cache
176
+ class BqaLMCache:
177
+ """Decoding state. `kv`: K/V of the attention layers (post norm and RoPE). The recurrent layers' states live in
178
+ their libraries' own containers, created on first use: mamba_ssm's InferenceParams (Mamba2 conv / SSM states,
179
+ updated in place) and fla's FLACache (GatedDeltaNet / KDA recurrent and conv states, indexed by layer)."""
180
+ is_compileable = False
181
+
182
+ def __init__(self, batch_size: int, max_seqlen: int):
183
+ self.kv = {}
184
+ self.batch_size = batch_size
185
+ self.max_seqlen = max_seqlen
186
+ self.seen = 0
187
+ self._mamba = None
188
+ self._fla = None
189
+
190
+ @property
191
+ def mamba_params(self):
192
+ if self._mamba is None:
193
+ try:
194
+ from mamba_ssm.utils.generation import InferenceParams
195
+ except ImportError as e:
196
+ raise ImportError("this checkpoint's Mamba2 layers need `pip install mamba-ssm causal-conv1d`") from e
197
+ self._mamba = InferenceParams(max_seqlen=self.max_seqlen, max_batch_size=self.batch_size)
198
+ self._mamba.seqlen_offset = self.seen
199
+ return self._mamba
200
+
201
+ @property
202
+ def fla(self):
203
+ if self._fla is None:
204
+ try:
205
+ from fla.models.utils import FLACache
206
+ except ImportError as e:
207
+ raise ImportError("this checkpoint's GatedDeltaNet / KDA layers need `pip install flash-linear-attention`") from e
208
+ self._fla = FLACache()
209
+ return self._fla
210
+
211
+ def get_seq_length(self, layer_idx: int = 0) -> int:
212
+ return self.seen
213
+
214
+ def advance(self, n: int) -> None:
215
+ self.seen += n
216
+ if self._mamba is not None:
217
+ self._mamba.seqlen_offset += n
218
+
219
+
220
+ # ---------------------------------------------------------------------------------------------------- attention
221
+ def _pick_flash():
222
+ """flash-attn's flash_attn_func if installed (accepted only if it takes `dropout_p`), else None."""
223
+ import inspect
224
+ for mod in ("flash_attn_interface", "flash_attn"):
225
+ try:
226
+ fn = getattr(__import__(mod, fromlist=["flash_attn_func"]), "flash_attn_func")
227
+ if "dropout_p" in inspect.signature(fn).parameters:
228
+ return fn
229
+ except Exception:
230
+ continue
231
+ return None
232
+
233
+
234
+ _FLASH_FN = _pick_flash()
235
+
236
+
237
+ class GQAAttention(nn.Module):
238
+ """Grouped-query attention: qk-norm before (partial) RoPE, optional sigmoid output gate, SDPA or flash-attn."""
239
+
240
+ def __init__(self, cfg: ModelConfig, layer_idx: int):
241
+ super().__init__()
242
+ self.layer_idx = layer_idx
243
+ self.n_heads = cfg.n_heads
244
+ self.n_kv_heads = cfg.n_kv_heads
245
+ self.n_rep = cfg.n_rep
246
+ self.head_dim = cfg.head_dim
247
+ self.attn_dropout = cfg.attn_dropout
248
+ impl = "flash" if (cfg.attn_impl == "auto" and _FLASH_FN is not None) else cfg.attn_impl
249
+ if impl not in ("flash", "sdpa", "auto"):
250
+ raise ValueError(f"this checkpoint's remote code supports attn_impl sdpa | flash, got {impl!r}")
251
+ if impl == "flash" and _FLASH_FN is None:
252
+ raise RuntimeError("attn_impl='flash' but flash-attn is not importable; use attn_implementation='sdpa'")
253
+ self.impl = "sdpa" if impl == "auto" else impl
254
+ self.attn_mscale = yarn_mscale(cfg.rope_scaling)
255
+ self.attn_output_gate = cfg.attn_output_gate
256
+ self.q_proj = nn.Linear(cfg.d_model, q_proj_out_features(cfg), bias=False)
257
+ self.k_proj = nn.Linear(cfg.d_model, cfg.kv_dim, bias=False)
258
+ self.v_proj = nn.Linear(cfg.d_model, cfg.kv_dim, bias=False)
259
+ self.o_proj = nn.Linear(cfg.q_dim, cfg.d_model, bias=False)
260
+ self.qk_norm = cfg.qk_norm
261
+ if self.qk_norm:
262
+ self.q_norm = RMSNorm(self.head_dim, cfg.norm_eps, cfg.rms_norm_in_fp32)
263
+ self.k_norm = RMSNorm(self.head_dim, cfg.norm_eps, cfg.rms_norm_in_fp32)
264
+
265
+ def forward(self, x, cos, sin, attention_mask=None, cache=None):
266
+ B, T, _ = x.shape
267
+ q, gate = split_q_gate(self.q_proj(x), self.n_heads, self.head_dim, self.attn_output_gate)
268
+ q = q.transpose(1, 2) # [B, H, T, D]
269
+ k = self.k_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2) # [B, Hkv, T, D]
270
+ v = self.v_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2)
271
+ if self.qk_norm:
272
+ q, k = self.q_norm(q), self.k_norm(k)
273
+ q, k = apply_rotary(q, k, cos, sin)
274
+ if self.attn_mscale != 1.0:
275
+ q = q * (self.attn_mscale * self.attn_mscale)
276
+ drop = self.attn_dropout if self.training else 0.0
277
+ past = None
278
+ if cache is not None:
279
+ past = cache.kv.get(self.layer_idx)
280
+ if past is not None:
281
+ k = torch.cat([past[0], k], dim=2)
282
+ v = torch.cat([past[1], v], dim=2)
283
+ cache.kv[self.layer_idx] = (k, v)
284
+ if past is not None: # decoding: the T new queries sit at the end of the S cached positions
285
+ S = k.shape[2]
286
+ gqa = self.n_rep > 1
287
+ if attention_mask is None and T == 1:
288
+ out = F.scaled_dot_product_attention(q, k, v, enable_gqa=gqa)
289
+ else:
290
+ pos_q = torch.arange(S - T, S, device=q.device)
291
+ keep = (torch.arange(S, device=q.device)[None, :] <= pos_q[:, None])[None, None]
292
+ if attention_mask is not None:
293
+ keep = keep & attention_mask.bool()[:, None, None, -S:]
294
+ bias = torch.zeros(keep.shape, dtype=q.dtype, device=q.device)
295
+ bias.masked_fill_(~keep, torch.finfo(q.dtype).min)
296
+ out = F.scaled_dot_product_attention(q, k, v, attn_mask=bias, enable_gqa=gqa)
297
+ out = out.transpose(1, 2).reshape(B, T, self.n_heads * self.head_dim)
298
+ elif self.impl == "flash" and attention_mask is None:
299
+ out = _FLASH_FN(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), dropout_p=drop, causal=True)
300
+ if isinstance(out, tuple):
301
+ out = out[0]
302
+ out = out.reshape(B, T, self.n_heads * self.head_dim)
303
+ else:
304
+ gqa = self.n_rep > 1
305
+ if attention_mask is None:
306
+ out = F.scaled_dot_product_attention(q, k, v, is_causal=True, dropout_p=drop, enable_gqa=gqa)
307
+ else:
308
+ bias = self._build_additive_mask(attention_mask, T, q.dtype, q.device)
309
+ out = F.scaled_dot_product_attention(q, k, v, attn_mask=bias, dropout_p=drop, enable_gqa=gqa)
310
+ out = out.transpose(1, 2).reshape(B, T, self.n_heads * self.head_dim)
311
+ return self.o_proj(apply_output_gate(out, gate))
312
+
313
+ @staticmethod
314
+ def _build_additive_mask(attention_mask, T, dtype, device):
315
+ causal = torch.tril(torch.ones(T, T, dtype=torch.bool, device=device))
316
+ keep = attention_mask.bool()[:, None, None, :] & causal[None, None]
317
+ bias = torch.zeros(keep.shape, dtype=dtype, device=device)
318
+ bias.masked_fill_(~keep, torch.finfo(dtype).min)
319
+ return bias
320
+
321
+
322
+ # ------------------------------------------------------------------------------------------------------- Mamba2
323
+ class Mamba2Mixer(nn.Module):
324
+ """mamba_ssm.Mamba2 on its fused kernel path (conv1d + SSD scan); RoPE inputs are accepted and ignored."""
325
+
326
+ def __init__(self, cfg: ModelConfig, layer_idx: int):
327
+ super().__init__()
328
+ try:
329
+ from mamba_ssm import Mamba2
330
+ import causal_conv1d # noqa: F401 (the fused kernel path needs it)
331
+ except ImportError as e:
332
+ raise ImportError("this checkpoint's Mamba2 layers need `pip install mamba-ssm causal-conv1d`") from e
333
+ headdim = int(cfg.mamba2_headdim) if getattr(cfg, "mamba2_headdim", None) else cfg.head_dim
334
+ self.mamba = Mamba2(
335
+ d_model=cfg.d_model,
336
+ headdim=headdim,
337
+ d_state=int(getattr(cfg, "mamba2_d_state", 128)),
338
+ expand=int(getattr(cfg, "mamba2_expand", 2)),
339
+ ngroups=int(getattr(cfg, "mamba2_ngroups", 1)),
340
+ chunk_size=int(getattr(cfg, "mamba2_chunk_size", 256)),
341
+ layer_idx=layer_idx,
342
+ )
343
+
344
+ def forward(self, x, cos=None, sin=None, attention_mask=None, cache=None):
345
+ return self.mamba(x, inference_params=cache.mamba_params if cache is not None else None)
346
+
347
+
348
+ # ------------------------------------------------------------------------------------------------ GatedDeltaNet
349
+ class GatedDeltaNetMixer(nn.Module):
350
+ """fla's GatedDeltaNet (chunk mode, gate, short convolution); RoPE inputs are accepted and ignored."""
351
+
352
+ def __init__(self, cfg: ModelConfig, layer_idx: int):
353
+ super().__init__()
354
+ try:
355
+ from fla.layers import GatedDeltaNet
356
+ except ImportError as e:
357
+ raise ImportError("this checkpoint's GatedDeltaNet layers need `pip install flash-linear-attention`") from e
358
+ head_dim = cfg.gdn_head_dim or cfg.head_dim
359
+ num_heads = cfg.gdn_num_heads or max(1, cfg.d_model // head_dim)
360
+ extra = {}
361
+ if getattr(cfg, "gdn_num_v_heads", None): # omitted when unset: passing None is not equivalent in older fla
362
+ extra["num_v_heads"] = cfg.gdn_num_v_heads
363
+ self.gdn = GatedDeltaNet(hidden_size=cfg.d_model, head_dim=head_dim, num_heads=num_heads, **extra,
364
+ mode="chunk", use_gate=True, use_short_conv=True, layer_idx=layer_idx)
365
+
366
+ def forward(self, x, cos=None, sin=None, attention_mask=None, cache=None):
367
+ out = self.gdn(x) if cache is None else self.gdn(x, past_key_values=cache.fla, use_cache=True)
368
+ return out[0] if isinstance(out, tuple) else out
369
+
370
+
371
+ # ---------------------------------------------------------------------------------------------------------- KDA
372
+ class KDAMixer(nn.Module):
373
+ """fla's Kimi Delta Attention (chunk mode, short convolution, per-channel decay); RoPE inputs are accepted and
374
+ ignored. head_dim / num_heads fall back to the gdn_* fields, as in training."""
375
+
376
+ def __init__(self, cfg: ModelConfig, layer_idx: int):
377
+ super().__init__()
378
+ try:
379
+ from fla.layers.kda import KimiDeltaAttention
380
+ except ImportError as e:
381
+ raise ImportError("this checkpoint's KDA layers need `pip install flash-linear-attention`") from e
382
+ head_dim = cfg.kda_head_dim or cfg.gdn_head_dim or cfg.head_dim
383
+ num_heads = cfg.kda_num_heads or cfg.gdn_num_heads or max(1, cfg.d_model // head_dim)
384
+ extra = {}
385
+ if getattr(cfg, "kda_num_v_heads", None):
386
+ extra["num_v_heads"] = int(cfg.kda_num_v_heads)
387
+ if getattr(cfg, "kda_lower_bound", None) is not None:
388
+ extra["lower_bound"] = float(cfg.kda_lower_bound)
389
+ self.kda = KimiDeltaAttention(hidden_size=cfg.d_model, head_dim=head_dim, num_heads=num_heads,
390
+ expand_v=float(getattr(cfg, "kda_expand_v", 1.0)), mode="chunk",
391
+ use_short_conv=True, conv_size=int(getattr(cfg, "kda_conv_size", 4)),
392
+ allow_neg_eigval=bool(getattr(cfg, "kda_allow_neg_eigval", False)),
393
+ safe_gate=bool(getattr(cfg, "kda_safe_gate", False)), layer_idx=layer_idx,
394
+ **extra)
395
+
396
+ def forward(self, x, cos=None, sin=None, attention_mask=None, cache=None):
397
+ out = self.kda(x) if cache is None else self.kda(x, past_key_values=cache.fla, use_cache=True)
398
+ return out[0] if isinstance(out, tuple) else out
399
+
400
+
401
+ # ---------------------------------------------------------------------------------------------------------- MLA
402
+ def _rope_leading(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, rd: int) -> torch.Tensor:
403
+ """Rotate the leading `rd` channels of x [B, h, T, D] (h = 1 or H) with cos / sin [B, T, rd]; the tail passes."""
404
+ c, s = cos.unsqueeze(1), sin.unsqueeze(1)
405
+ xr = x[..., :rd].float()
406
+ out = (xr * c + rotate_half(xr) * s).to(x.dtype)
407
+ return out if rd == x.shape[-1] else torch.cat([out, x[..., rd:]], dim=-1)
408
+
409
+
410
+ class MLAAttention(nn.Module):
411
+ """Multi-head latent attention at one head dim for q, k and v (the bf16 core the study's evaluation ran).
412
+
413
+ q = q_proj(x) per head [rope | nope]; [c ; k_r] = kv_a_proj(x); c = kv_a_norm(c); k_nope = k_proj(c);
414
+ v = v_proj(c); four sub-vector norms (qk_norm); k_r rotated once and shared by every head; MHA on the wire."""
415
+
416
+ def __init__(self, cfg: ModelConfig, layer_idx: int):
417
+ super().__init__()
418
+ self.layer_idx = layer_idx
419
+ self.n_heads = cfg.n_heads
420
+ self.head_dim = cfg.head_dim
421
+ self.rotary_dim = cfg.rotary_dim
422
+ self.nope_dim = cfg.head_dim - cfg.rotary_dim
423
+ self.kv_lora_rank = int(cfg.kv_lora_rank)
424
+ self.attn_output_gate = cfg.attn_output_gate
425
+ self.qk_norm = cfg.qk_norm
426
+ assert cfg.n_kv_heads == cfg.n_heads, "MLA is MHA on the wire: n_kv_heads must equal n_heads"
427
+ assert 0 < self.rotary_dim < self.head_dim, "MLA needs 0 < rotary_dim < head_dim"
428
+ self.attn_mscale = yarn_mscale(cfg.rope_scaling)
429
+ d, H, D, rd, dn, dc = cfg.d_model, self.n_heads, self.head_dim, self.rotary_dim, self.nope_dim, self.kv_lora_rank
430
+ self.q_proj = nn.Linear(d, q_proj_out_features(cfg), bias=False)
431
+ self.kv_a_proj = nn.Linear(d, dc + rd, bias=False)
432
+ self.kv_a_norm = RMSNorm(dc, cfg.norm_eps, cfg.rms_norm_in_fp32)
433
+ self.k_proj = nn.Linear(dc, H * dn, bias=False)
434
+ self.v_proj = nn.Linear(dc, H * D, bias=False)
435
+ self.o_proj = nn.Linear(H * D, d, bias=False)
436
+ if self.qk_norm:
437
+ self.q_rope_norm = RMSNorm(rd, cfg.norm_eps, cfg.rms_norm_in_fp32)
438
+ self.q_nope_norm = RMSNorm(dn, cfg.norm_eps, cfg.rms_norm_in_fp32)
439
+ self.k_rope_norm = RMSNorm(rd, cfg.norm_eps, cfg.rms_norm_in_fp32)
440
+ self.k_nope_norm = RMSNorm(dn, cfg.norm_eps, cfg.rms_norm_in_fp32)
441
+ impl = "flash" if (cfg.attn_impl == "auto" and _FLASH_FN is not None) else cfg.attn_impl
442
+ if impl == "flash" and _FLASH_FN is None:
443
+ raise RuntimeError("attn_impl='flash' but flash-attn is not importable; use attn_implementation='sdpa'")
444
+ self.impl = impl if impl == "flash" else "sdpa"
445
+
446
+ def _qkv(self, x, cos, sin):
447
+ B, T, _ = x.shape
448
+ H, D, rd, dn, dc = self.n_heads, self.head_dim, self.rotary_dim, self.nope_dim, self.kv_lora_rank
449
+ q, gate = split_q_gate(self.q_proj(x), H, D, self.attn_output_gate)
450
+ q = q.transpose(1, 2) # [B, H, T, D]
451
+ ckr = self.kv_a_proj(x)
452
+ c, k_r = ckr[..., :dc], ckr[..., dc:] # [B, T, dc], [B, T, rd]
453
+ c = self.kv_a_norm(c)
454
+ k_n = self.k_proj(c).view(B, T, H, dn).transpose(1, 2) # [B, H, T, dn]
455
+ v = self.v_proj(c).view(B, T, H, D).transpose(1, 2) # [B, H, T, D]
456
+ k_r = k_r.unsqueeze(1) # [B, 1, T, rd]
457
+ if self.qk_norm:
458
+ q = torch.cat([self.q_rope_norm(q[..., :rd]), self.q_nope_norm(q[..., rd:])], dim=-1)
459
+ k_r = self.k_rope_norm(k_r)
460
+ k_n = self.k_nope_norm(k_n)
461
+ q = _rope_leading(q, cos, sin, rd)
462
+ if self.attn_mscale != 1.0:
463
+ q = q * (self.attn_mscale * self.attn_mscale)
464
+ k_r = _rope_leading(k_r, cos, sin, rd)
465
+ k = torch.cat([k_r.expand(B, H, T, rd), k_n], dim=-1)
466
+ return q, k, v, gate
467
+
468
+ def forward(self, x, cos, sin, attention_mask=None, cache=None):
469
+ B, T, _ = x.shape
470
+ q, k, v, gate = self._qkv(x, cos, sin)
471
+ past = None
472
+ if cache is not None:
473
+ past = cache.kv.get(self.layer_idx)
474
+ if past is not None:
475
+ k = torch.cat([past[0], k], dim=2)
476
+ v = torch.cat([past[1], v], dim=2)
477
+ cache.kv[self.layer_idx] = (k, v)
478
+ S = k.shape[2]
479
+ if attention_mask is None and past is None and self.impl == "flash":
480
+ out = _FLASH_FN(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), causal=True)
481
+ if isinstance(out, tuple):
482
+ out = out[0]
483
+ out = out.reshape(B, T, self.n_heads * self.head_dim)
484
+ else:
485
+ if attention_mask is None and past is None:
486
+ out = F.scaled_dot_product_attention(q, k, v, is_causal=True)
487
+ elif attention_mask is None and T == 1:
488
+ out = F.scaled_dot_product_attention(q, k, v)
489
+ else:
490
+ pos_q = torch.arange(S - T, S, device=q.device)
491
+ keep = (torch.arange(S, device=q.device)[None, :] <= pos_q[:, None])[None, None]
492
+ if attention_mask is not None:
493
+ keep = keep & attention_mask.bool()[:, None, None, -S:]
494
+ bias = torch.zeros(keep.shape, dtype=q.dtype, device=q.device)
495
+ bias.masked_fill_(~keep, torch.finfo(q.dtype).min)
496
+ out = F.scaled_dot_product_attention(q, k, v, attn_mask=bias)
497
+ out = out.transpose(1, 2).reshape(B, T, self.n_heads * self.head_dim)
498
+ return self.o_proj(apply_output_gate(out, gate))
499
+
500
+
501
+ MIXER_REGISTRY = {"gqa": GQAAttention, "mla": MLAAttention, "mamba2": Mamba2Mixer, "gated_deltanet": GatedDeltaNetMixer,
502
+ "kda": KDAMixer}
503
+
504
+
505
+ def build_mixer(name: str, cfg: ModelConfig, layer_idx: int) -> nn.Module:
506
+ if name not in MIXER_REGISTRY:
507
+ raise ValueError(f"mixer {name!r} is not shipped with this checkpoint's code (have {sorted(MIXER_REGISTRY)})")
508
+ return MIXER_REGISTRY[name](cfg, layer_idx)
509
+
510
+
511
+ # ------------------------------------------------------------------------------------------------------ backbone
512
+ class DecoderBlock(nn.Module):
513
+ """Pre-norm block: x += mixer(norm(x)); x += mlp(norm(x))."""
514
+
515
+ def __init__(self, cfg: ModelConfig, layer_idx: int):
516
+ super().__init__()
517
+ self.layer_idx = layer_idx
518
+ self.input_layernorm = RMSNorm(cfg.d_model, cfg.norm_eps, cfg.rms_norm_in_fp32)
519
+ self.mixer = build_mixer(cfg.mixer_for_layer(layer_idx), cfg, layer_idx)
520
+ self.post_attention_layernorm = RMSNorm(cfg.d_model, cfg.norm_eps, cfg.rms_norm_in_fp32)
521
+ self.mlp = SwiGLUMLP(cfg.d_model, cfg.intermediate_size, cfg.resid_dropout)
522
+ self.resid_drop = nn.Dropout(cfg.resid_dropout) if cfg.resid_dropout > 0 else nn.Identity()
523
+
524
+ def forward(self, x, cos, sin, attention_mask=None, cache=None):
525
+ x = x + self.resid_drop(self.mixer(self.input_layernorm(x), cos, sin, attention_mask, cache))
526
+ x = x + self.resid_drop(self.mlp(self.post_attention_layernorm(x)))
527
+ return x
528
+
529
+
530
+ class Transformer(nn.Module):
531
+ def __init__(self, cfg: ModelConfig):
532
+ super().__init__()
533
+ self.cfg = cfg
534
+ self.embed_tokens = nn.Embedding(cfg.vocab_size, cfg.d_model)
535
+ self.layers = nn.ModuleList([DecoderBlock(cfg, i) for i in range(cfg.n_layers)])
536
+ self.norm = RMSNorm(cfg.d_model, cfg.norm_eps, cfg.rms_norm_in_fp32)
537
+ self.rotary = RotaryEmbedding(cfg.head_dim, cfg.max_seq_len, cfg.rope_theta,
538
+ rope_scaling=cfg.rope_scaling, rotary_dim=cfg.rotary_dim)
539
+
540
+ def forward(self, input_ids, position_ids, attention_mask=None, cache=None):
541
+ if attention_mask is not None and bool(attention_mask.all()):
542
+ attention_mask = None # no padding: take the maskless kernels, the same path as a plain forward
543
+ h = self.embed_tokens(input_ids)
544
+ cos, sin = self.rotary(position_ids)
545
+ if getattr(self.cfg, "nope", False):
546
+ cos, sin = torch.ones_like(cos), torch.zeros_like(sin)
547
+ cos, sin = cos.to(h.dtype), sin.to(h.dtype)
548
+ for layer in self.layers:
549
+ h = layer(h, cos, sin, attention_mask, cache)
550
+ if cache is not None:
551
+ cache.advance(input_ids.shape[1])
552
+ return self.norm(h)
553
+
554
+
555
+ # ------------------------------------------------------------------------------------------------ HF causal LM
556
+ _ATTN_MAP = {"flash_attention_2": "flash", "sdpa": "sdpa", "eager": "sdpa"}
557
+
558
+
559
+ @dataclass
560
+ class BqaLMCausalLMOutput(ModelOutput):
561
+ loss: torch.FloatTensor | None = None
562
+ logits: torch.FloatTensor | None = None
563
+ cache_params: BqaLMCache | None = None
564
+
565
+
566
+ class BqaLMForCausalLM(PreTrainedModel, GenerationMixin):
567
+ config_class = BqaLMConfig
568
+ base_model_prefix = "model"
569
+ supports_gradient_checkpointing = True
570
+ _supports_flash_attn = True
571
+ _supports_flash_attn_2 = True
572
+ _supports_sdpa = True
573
+ _supports_attention_backend = True
574
+ _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"}
575
+ _is_stateful = True # the Mamba2 state cannot be rewound, so no assisted generation
576
+
577
+ @classmethod
578
+ def _supports_default_dynamic_cache(cls) -> bool:
579
+ return False # the model keeps its own state in `cache_params` (BqaLMCache)
580
+
581
+ def __init__(self, config: BqaLMConfig):
582
+ super().__init__(config)
583
+ mc = config.to_model_config()
584
+ hf_impl = getattr(config, "_attn_implementation", None)
585
+ if hf_impl in _ATTN_MAP:
586
+ mc.attn_impl = _ATTN_MAP[hf_impl]
587
+ self._mc = mc
588
+ self.model = Transformer(mc)
589
+ self.lm_head = nn.Linear(mc.d_model, mc.vocab_size, bias=False)
590
+ if mc.tie_embeddings:
591
+ self.lm_head.weight = self.model.embed_tokens.weight
592
+ self.post_init()
593
+
594
+ def get_input_embeddings(self):
595
+ return self.model.embed_tokens
596
+
597
+ def set_input_embeddings(self, value):
598
+ self.model.embed_tokens = value
599
+
600
+ def get_output_embeddings(self):
601
+ return self.lm_head
602
+
603
+ def set_output_embeddings(self, value):
604
+ self.lm_head = value
605
+
606
+ def forward(self, input_ids=None, attention_mask=None, position_ids=None, labels=None,
607
+ cache_params=None, use_cache=None, past_key_values=None, output_attentions=None,
608
+ output_hidden_states=None, return_dict=None, **kwargs):
609
+ B, T = input_ids.shape
610
+ if use_cache and cache_params is None:
611
+ cache_params = BqaLMCache(B, self._mc.max_seq_len)
612
+ past = cache_params.seen if cache_params is not None else 0
613
+ if position_ids is None:
614
+ position_ids = torch.arange(past, past + T, device=input_ids.device).unsqueeze(0).expand(B, -1)
615
+ hidden = self.model(input_ids, position_ids, attention_mask, cache=cache_params)
616
+ logits = self.lm_head(hidden)
617
+ loss = None
618
+ if labels is not None:
619
+ sl = logits[:, :-1, :].contiguous().float()
620
+ lb = labels[:, 1:].contiguous()
621
+ loss = F.cross_entropy(sl.view(-1, sl.size(-1)), lb.view(-1), ignore_index=-100)
622
+ return BqaLMCausalLMOutput(loss=loss, logits=logits, cache_params=cache_params)
623
+
624
+ def prepare_inputs_for_generation(self, input_ids, attention_mask=None, cache_params=None, use_cache=None,
625
+ **kwargs):
626
+ # with a cache only the newest token is fed; use_cache=False recomputes the whole prefix every step
627
+ if cache_params is not None and cache_params.seen > 0:
628
+ input_ids = input_ids[:, -1:]
629
+ return {"input_ids": input_ids, "attention_mask": attention_mask, "cache_params": cache_params,
630
+ "use_cache": use_cache}
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer_config.json ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": null,
3
+ "backend": "tokenizers",
4
+ "bos_token": "<s>",
5
+ "clean_up_tokenization_spaces": false,
6
+ "eos_token": "</s>",
7
+ "is_local": true,
8
+ "local_files_only": true,
9
+ "model_max_length": 1000000000000000019884624838656,
10
+ "pad_token": null,
11
+ "padding_side": "right",
12
+ "sp_model_kwargs": {},
13
+ "tokenizer_class": "LlamaTokenizer",
14
+ "unk_token": "<unk>",
15
+ "use_default_system_prompt": false
16
+ }