korbip commited on
Commit
92db682
·
verified ·
1 Parent(s): e5b6ee6

Add ComplexKDA 1.3B/100BT FineWeb-Edu model

Browse files
README.md ADDED
@@ -0,0 +1,100 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ library_name: transformers
3
+ license: mit
4
+ pipeline_tag: text-generation
5
+ tags:
6
+ - complex-kda
7
+ - linear-attention
8
+ - kimi-delta-attention
9
+ ---
10
+
11
+ # openeurollm/kda-sigmoid-hybrid-1.3B-100B
12
+
13
+ A hybrid ComplexKDA language model (1.36B parameters).
14
+
15
+ ComplexKDA is Kimi Delta Attention with a **signed decay gate**: the per-channel
16
+ decay `alpha` is allowed to take either sign, `alpha in [-1, 1]`, instead of
17
+ being confined to `(0, 1]`. That is the one-dimensional real case of a complex
18
+ eigenvalue, so a channel can oscillate rather than only forget. The magnitude is
19
+ carried in log space exactly as KDA carries it; the `+-1` part is carried as a
20
+ running product pushed onto the queries and keys, so the recurrence the kernels
21
+ run is still the unsigned one.
22
+
23
+ This checkpoint's decay gate is **unsigned sigmoid (the KDA baseline: alpha in (0, 1])**.
24
+
25
+ ## Architecture
26
+
27
+ - 24 layers, hidden size 2048, MLP 5312 (SwiGLU)
28
+ - 16 heads of dimension 128, short convolution of width 4
29
+ - vocabulary 32000, trained at context 4096
30
+ - embeddings untied
31
+ - attention at layers [3, 7, 11, 15, 19, 23] (gated, NoPE), linear everywhere else
32
+
33
+ ## Tokenizer, and how to start a prompt
34
+
35
+ The bundled tokenizer is configured the way the training corpus was encoded:
36
+ **no BOS is prepended**, documents were terminated with the EOS token, and
37
+ `model_max_length` is this model's trained context. The upstream tokenizer
38
+ repository's own defaults differ on both points, so encode through the
39
+ tokenizer shipped here rather than re-fetching it by name.
40
+
41
+ **To condition on the start of a document, prefix the EOS token, not a BOS.**
42
+ Training packed documents as `[text..., EOS]`, so the token preceding any
43
+ document's first token is always the EOS; a BOS was never seen at any position.
44
+ Measured on held-out documents at 1.3B, against a bare prompt:
45
+
46
+ | prefix | mean NLL | first 8 tokens |
47
+ |---|---|---|
48
+ | none | 2.0104 | 3.2947 |
49
+ | **EOS** | **1.9960** | **2.9606** |
50
+ | BOS | 2.0150 | 3.3768 |
51
+
52
+ So the EOS is worth ~0.33 nats on the opening tokens and the BOS costs ~0.08.
53
+
54
+ **Only at a document start.** The prefix is a "a new document begins here"
55
+ signal, and mid-document it is a false one — on continuations the same EOS
56
+ prefix *costs* 0.34 nats over the first 8 tokens (2.19 → 2.22 overall). Prefix
57
+ it when you mean a fresh document; leave a continuation bare. Nothing in the
58
+ weights or the tokenizer can tell these apart, which is why this is the
59
+ caller's decision.
60
+
61
+ `bos_token` is remapped to `</s>` here, so a pipeline that asks for "the BOS"
62
+ (lm-eval's `--add_bos_token`, a generic wrapper) gets the separator rather than
63
+ the unused `<s>`. `add_bos_token` remains `False`, so the default is still a
64
+ bare prompt — which is the convention every number quoted for these models was
65
+ measured under.
66
+
67
+ If you are coming from `fla-hub` checkpoints, note that they use the opposite
68
+ convention -- trained as `[BOS, text...]` with the BOS as the document
69
+ separator -- so the habit does not carry over.
70
+
71
+ ## Usage
72
+
73
+ The bundled `modeling_complex_kda.py` is **standalone**: `torch` and
74
+ `transformers` are all it needs.
75
+
76
+ ```python
77
+ from transformers import AutoModelForCausalLM, AutoTokenizer
78
+
79
+ tok = AutoTokenizer.from_pretrained("openeurollm/kda-sigmoid-hybrid-1.3B-100B")
80
+ model = AutoModelForCausalLM.from_pretrained(
81
+ "openeurollm/kda-sigmoid-hybrid-1.3B-100B", trust_remote_code=True, dtype="bfloat16")
82
+ ```
83
+
84
+ For the Triton kernels these models were trained with -- much faster, and the
85
+ exact code path of the training runs -- install the fork:
86
+
87
+ ```bash
88
+ pip install git+https://github.com/automl/ComplexKDA
89
+ ```
90
+
91
+ It is picked up automatically when importable. `COMPLEX_KDA_BACKEND=torch`
92
+ forces the portable path; `=kernel` makes a missing fork an error instead of a
93
+ silent fallback.
94
+
95
+ ## Provenance
96
+
97
+ Converted from the training checkpoint with `lm_scaling/hf_release/convert_to_hub.py`.
98
+ The conversion is metadata only -- the weight file is the exporter's own, byte
99
+ for byte -- and the bundled implementation is checked against the reference
100
+ implementation the runs used.
config.json ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "ComplexKDAForCausalLM"
4
+ ],
5
+ "model_type": "complex_kda",
6
+ "auto_map": {
7
+ "AutoConfig": "configuration_complex_kda.ComplexKDAConfig",
8
+ "AutoModel": "modeling_complex_kda.ComplexKDAModel",
9
+ "AutoModelForCausalLM": "modeling_complex_kda.ComplexKDAForCausalLM"
10
+ },
11
+ "attn_mode": "chunk",
12
+ "hidden_size": 2048,
13
+ "num_hidden_layers": 24,
14
+ "num_heads": 16,
15
+ "head_dim": 128,
16
+ "intermediate_size": 5312,
17
+ "hidden_ratio": null,
18
+ "hidden_act": "swish",
19
+ "vocab_size": 32000,
20
+ "max_position_embeddings": 4096,
21
+ "tie_word_embeddings": false,
22
+ "norm_eps": 1e-06,
23
+ "conv_size": 4,
24
+ "use_short_conv": true,
25
+ "fuse_norm": true,
26
+ "fuse_swiglu": true,
27
+ "fuse_cross_entropy": true,
28
+ "initializer_range": 0.02,
29
+ "allow_neg_eigval": false,
30
+ "lower_bound": -5.0,
31
+ "expand_v": 1.0,
32
+ "gate": "sigmoid",
33
+ "gate_init_style": "shipped",
34
+ "output_gate": "lowrank",
35
+ "attn": {
36
+ "layers": [
37
+ 3,
38
+ 7,
39
+ 11,
40
+ 15,
41
+ 19,
42
+ 23
43
+ ],
44
+ "num_heads": 16,
45
+ "num_kv_heads": 16,
46
+ "qkv_bias": false,
47
+ "qk_norm": false,
48
+ "output_gate": true,
49
+ "use_rope": false,
50
+ "rope_theta": 10000.0,
51
+ "window_size": null
52
+ },
53
+ "drop_silu": true,
54
+ "use_cache": true
55
+ }
configuration_complex_kda.py ADDED
@@ -0,0 +1,246 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2026, the ComplexKDA authors.
2
+ # Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li (the parts derived
3
+ # from flash-linear-attention, MIT licensed).
4
+ #
5
+ # SPDX-License-Identifier: MIT
6
+ """Config for the ComplexKDA models published on the Hub.
7
+
8
+ STANDALONE ON PURPOSE. `fla` is not imported anywhere in this file, so the
9
+ config loads with nothing but `transformers` installed. The hybrid-attention
10
+ normalisation that fla keeps in `fla/models/hybrid.py` is vendored below for
11
+ the same reason -- a config that needed the fork to parse would make every
12
+ error message about a missing package rather than about the model.
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ import math
18
+
19
+ from transformers.configuration_utils import PretrainedConfig
20
+
21
+ __all__ = ["ComplexKDAConfig"]
22
+
23
+
24
+ # ---------------------------------------------------------------------------
25
+ # hybrid attention spec (vendored from fla/models/hybrid.py)
26
+ #
27
+ # A hybrid arm replaces the linear mixer with full attention at the layers
28
+ # named in `attn["layers"]`. The dict is stored in config.json, so it is
29
+ # validated on the way in: an out-of-range or duplicated layer index would
30
+ # otherwise build a model whose `state_dict` silently disagrees with the
31
+ # checkpoint at exactly those layers.
32
+ # ---------------------------------------------------------------------------
33
+
34
+ def _spec_context(spec_index: int | None) -> str:
35
+ return "attn specification" if spec_index is None else f"attn specification at index {spec_index}"
36
+
37
+
38
+ def _positive_int(value: object, *, field: str, context: str) -> int:
39
+ if isinstance(value, bool) or not isinstance(value, int) or value <= 0:
40
+ raise ValueError(f"{context} field {field!r} must be a positive integer; got {value!r}")
41
+ return value
42
+
43
+
44
+ def _normalize_spec(spec: dict, *, num_hidden_layers: int, spec_index: int | None,
45
+ assigned_layers: dict) -> dict:
46
+ context = _spec_context(spec_index)
47
+ normalized = dict(spec)
48
+
49
+ for field in ("layers", "num_heads"):
50
+ if field not in normalized:
51
+ raise ValueError(f"{context} field {field!r} is required; got <missing>")
52
+
53
+ layers = normalized["layers"]
54
+ if not isinstance(layers, (list, tuple)):
55
+ raise ValueError(
56
+ f"{context} field 'layers' must be a list or tuple of integer layer indices; got {layers!r}")
57
+
58
+ normalized_layers, seen = [], set()
59
+ for layer_idx in layers:
60
+ if isinstance(layer_idx, bool) or not isinstance(layer_idx, int):
61
+ raise ValueError(f"{context} field 'layers' must contain only integer layer indices; got {layer_idx!r}")
62
+ if layer_idx < 0 or layer_idx >= num_hidden_layers:
63
+ raise ValueError(
64
+ f"{context} field 'layers' contains out-of-range layer {layer_idx!r}; "
65
+ f"expected a value in [0, {num_hidden_layers})")
66
+ if layer_idx in seen:
67
+ raise ValueError(f"{context} field 'layers' contains duplicate layer {layer_idx!r}; got {layers!r}")
68
+ if layer_idx in assigned_layers:
69
+ raise ValueError(
70
+ f"{context} assigns conflicting layer {layer_idx!r}, which is already assigned by "
71
+ f"{_spec_context(assigned_layers[layer_idx])}")
72
+ seen.add(layer_idx)
73
+ assigned_layers[layer_idx] = spec_index
74
+ normalized_layers.append(layer_idx)
75
+
76
+ normalized["layers"] = normalized_layers
77
+ normalized["num_heads"] = _positive_int(normalized["num_heads"], field="num_heads", context=context)
78
+
79
+ num_kv_heads = normalized.get("num_kv_heads")
80
+ if num_kv_heads is None:
81
+ num_kv_heads = normalized["num_heads"]
82
+ normalized["num_kv_heads"] = _positive_int(num_kv_heads, field="num_kv_heads", context=context)
83
+
84
+ qkv_bias = normalized.get("qkv_bias", False)
85
+ if not isinstance(qkv_bias, bool):
86
+ raise ValueError(f"{context} field 'qkv_bias' must be a Boolean; got {qkv_bias!r}")
87
+ normalized["qkv_bias"] = qkv_bias
88
+
89
+ window_size = normalized.get("window_size")
90
+ if window_size is not None:
91
+ window_size = _positive_int(window_size, field="window_size", context=context)
92
+ normalized["window_size"] = window_size
93
+
94
+ rope_theta = normalized.get("rope_theta", 10000.0)
95
+ try:
96
+ ok = (not isinstance(rope_theta, bool) and isinstance(rope_theta, (int, float))
97
+ and math.isfinite(rope_theta) and rope_theta > 0)
98
+ except OverflowError:
99
+ ok = False
100
+ if not ok:
101
+ raise ValueError(f"{context} field 'rope_theta' must be positive and finite; got {rope_theta!r}")
102
+ normalized["rope_theta"] = rope_theta
103
+
104
+ return normalized
105
+
106
+
107
+ def normalize_hybrid_attention_config(attn, *, num_hidden_layers: int):
108
+ """Validate and normalise `attn`: None, one spec dict, or a list of them."""
109
+ if attn is None:
110
+ return None
111
+ if isinstance(num_hidden_layers, bool) or not isinstance(num_hidden_layers, int) or num_hidden_layers < 0:
112
+ raise ValueError(f"field 'num_hidden_layers' must be a non-negative integer; got {num_hidden_layers!r}")
113
+ if not isinstance(attn, (dict, list)):
114
+ raise ValueError(f"attn must be None, a dictionary, or a list of dictionaries; got {attn!r}")
115
+
116
+ is_single = isinstance(attn, dict)
117
+ specs = [attn] if is_single else attn
118
+ assigned: dict = {}
119
+ out = []
120
+ for i, spec in enumerate(specs):
121
+ idx = None if is_single else i
122
+ if not isinstance(spec, dict):
123
+ raise ValueError(f"{_spec_context(idx)} must be a dictionary; got {spec!r}")
124
+ out.append(_normalize_spec(spec, num_hidden_layers=num_hidden_layers,
125
+ spec_index=idx, assigned_layers=assigned))
126
+ return out[0] if is_single else out
127
+
128
+
129
+ def get_hybrid_attention_spec(attn, *, layer_idx: int):
130
+ """The normalised spec assigned to `layer_idx`, or None for a linear layer."""
131
+ if attn is None:
132
+ return None
133
+ for spec in ([attn] if isinstance(attn, dict) else attn):
134
+ if layer_idx in spec["layers"]:
135
+ return spec
136
+ return None
137
+
138
+
139
+ class ComplexKDAConfig(PretrainedConfig):
140
+ """Kimi Delta Attention with a *signed* (complex, i.e. Z_2-phased) decay gate.
141
+
142
+ The baseline arms set ``gate="sigmoid"`` with ``allow_neg_eigval=False``;
143
+ the ComplexKDA arms set ``gate="signed_sigmoid2"`` with
144
+ ``allow_neg_eigval=True``. Everything else is shared, which is what makes
145
+ the two comparable.
146
+ """
147
+
148
+ model_type = "complex_kda"
149
+ keys_to_ignore_at_inference = ["past_key_values"]
150
+
151
+ def __init__(
152
+ self,
153
+ attn_mode: str = "chunk",
154
+ hidden_size: int = 2048,
155
+ expand_v: float = 1.0,
156
+ use_short_conv: bool = True,
157
+ drop_silu: bool = False,
158
+ drop_key_silu: bool = False,
159
+ conv_silu: str = "qkv",
160
+ allow_neg_eigval: bool = True,
161
+ gate: str = "signed_sigmoid2",
162
+ gate_init_style: str = "shipped",
163
+ output_gate: str = "lowrank",
164
+ beta_init_style: str = "standard",
165
+ lower_bound: float = -5.0,
166
+ num_heads: int = 16,
167
+ num_v_heads: int | None = None,
168
+ head_dim: int = 128,
169
+ num_hidden_layers: int = 24,
170
+ norm_eps: float = 1e-6,
171
+ conv_size: int = 4,
172
+ attn: dict | list | None = None,
173
+ hidden_ratio: int | None = 4,
174
+ intermediate_size: int | None = None,
175
+ hidden_act: str = "swish",
176
+ max_position_embeddings: int = 4096,
177
+ initializer_range: float = 0.02,
178
+ vocab_size: int = 32000,
179
+ tie_word_embeddings: bool = False,
180
+ use_cache: bool = True,
181
+ pad_token_id: int | None = None,
182
+ bos_token_id: int = 1,
183
+ eos_token_id: int = 2,
184
+ # Kept so a config written by the training stack round-trips. They
185
+ # select fused kernels when the fla fork is installed and are ignored
186
+ # by the pure-torch path, which computes the same thing either way.
187
+ fuse_norm: bool = True,
188
+ fuse_swiglu: bool = True,
189
+ fuse_cross_entropy: bool = True,
190
+ use_l2warp: bool = False,
191
+ **kwargs,
192
+ ):
193
+ self.attn_mode = attn_mode
194
+ self.hidden_size = hidden_size
195
+ self.expand_v = expand_v
196
+ self.use_short_conv = use_short_conv
197
+ self.drop_silu = drop_silu
198
+ self.drop_key_silu = drop_key_silu
199
+ self.conv_silu = conv_silu
200
+ self.allow_neg_eigval = allow_neg_eigval
201
+ self.gate = gate
202
+ self.gate_init_style = gate_init_style
203
+ self.output_gate = output_gate
204
+ self.beta_init_style = beta_init_style
205
+ self.lower_bound = lower_bound
206
+ self.num_heads = num_heads
207
+ self.num_v_heads = num_v_heads
208
+ self.head_dim = head_dim
209
+ # `num_hidden_layers` must be set before `attn`: the layer-range check
210
+ # in the setter below reads it.
211
+ self.num_hidden_layers = num_hidden_layers
212
+ self.norm_eps = norm_eps
213
+ self.conv_size = conv_size
214
+ self.attn = attn
215
+
216
+ self.hidden_ratio = hidden_ratio
217
+ self.intermediate_size = intermediate_size
218
+ self.hidden_act = hidden_act
219
+ self.max_position_embeddings = max_position_embeddings
220
+ self.initializer_range = initializer_range
221
+ self.vocab_size = vocab_size
222
+ self.use_cache = use_cache
223
+
224
+ self.fuse_norm = fuse_norm
225
+ self.fuse_swiglu = fuse_swiglu
226
+ self.fuse_cross_entropy = fuse_cross_entropy
227
+ self.use_l2warp = use_l2warp
228
+
229
+ super().__init__(
230
+ pad_token_id=pad_token_id,
231
+ bos_token_id=bos_token_id,
232
+ eos_token_id=eos_token_id,
233
+ tie_word_embeddings=tie_word_embeddings,
234
+ **kwargs,
235
+ )
236
+
237
+ # `attn` is a property so that assigning it after construction -- which
238
+ # `PretrainedConfig.from_dict` does -- is validated too.
239
+ @property
240
+ def attn(self):
241
+ return self.__dict__.get("attn")
242
+
243
+ @attn.setter
244
+ def attn(self, value) -> None:
245
+ self.__dict__["attn"] = normalize_hybrid_attention_config(
246
+ value, num_hidden_layers=self.num_hidden_layers)
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:4d261e3881c258757411e47066f053cdbd9169fc605114446e231de045fa4604
3
+ size 2724569536
modeling_complex_kda.py ADDED
@@ -0,0 +1,1383 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2026, the ComplexKDA authors.
2
+ # Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li (the parts derived
3
+ # from flash-linear-attention, MIT licensed).
4
+ #
5
+ # SPDX-License-Identifier: MIT
6
+ """ComplexKDA -- Kimi Delta Attention with a signed (Z_2-phased) decay gate.
7
+
8
+ STANDALONE. This file needs only `torch` and `transformers`. It carries its own
9
+ implementations of everything the model is made of -- the short convolution,
10
+ the RMS norms, the SwiGLU MLP, the attention layers of the hybrid arms, and the
11
+ gated-delta recurrence itself -- so a checkpoint loads and runs with nothing
12
+ else installed.
13
+
14
+ IT GOES FASTER WITH THE FORK. When `fla` from
15
+
16
+ https://github.com/automl/ComplexKDA
17
+
18
+ is importable, the Triton kernels and fused modules it ships are used instead,
19
+ and the model is then running exactly the code the checkpoints were trained
20
+ through. Detection is by capability, not by name: upstream flash-linear-attention
21
+ also provides `chunk_kda`, but without the `sign` argument the signed gate needs,
22
+ so the signature is inspected rather than trusted. Set the environment variable
23
+ `COMPLEX_KDA_BACKEND=torch` to force the pure-torch path (useful for debugging a
24
+ numerical difference), or `=kernel` to make a missing fork an error rather than a
25
+ silent fallback.
26
+
27
+ WHAT THE SIGNED GATE IS. A gated-delta layer carries a per-channel decay
28
+ `alpha`; KDA, like every gated linear attention before it, confines it to
29
+ `(0, 1]`. ComplexKDA lets it take either sign, `alpha in [-1, 1]`, which is the
30
+ one-dimensional real case of a complex eigenvalue -- a channel can now oscillate
31
+ rather than only forget. The magnitude is carried in log space exactly as
32
+ before, and the `+-1` part is carried separately as a running product (the
33
+ "gauge") pushed onto q and k, so the recurrence the kernels run is still the
34
+ unsigned one. `running_sign` below is that product, and `ungauge_state` takes it
35
+ back off the state at a chunk boundary so a cached state is the real one.
36
+ """
37
+
38
+ from __future__ import annotations
39
+
40
+ import math
41
+ import os
42
+ import warnings
43
+ from typing import Any
44
+
45
+ import torch
46
+ import torch.nn as nn
47
+ import torch.nn.functional as F
48
+ from transformers.generation import GenerationMixin
49
+ from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast
50
+ from transformers.modeling_utils import PreTrainedModel
51
+ from transformers.utils import logging
52
+
53
+ try: # packaged next to the weights (the Hub layout)
54
+ from .configuration_complex_kda import ComplexKDAConfig, get_hybrid_attention_spec
55
+ except ImportError: # imported as a loose file
56
+ from configuration_complex_kda import ComplexKDAConfig, get_hybrid_attention_spec
57
+
58
+ logger = logging.get_logger(__name__)
59
+
60
+ __all__ = [
61
+ "ComplexKDACache",
62
+ "ComplexKDAForCausalLM",
63
+ "ComplexKDAModel",
64
+ "ComplexKDAPreTrainedModel",
65
+ "ComplexKimiDeltaAttention",
66
+ ]
67
+
68
+
69
+ # ===========================================================================
70
+ # optional fast path
71
+ # ===========================================================================
72
+
73
+ def _detect_fla():
74
+ """(chunk_kda, fused_recurrent_kda) from the fork, or (None, None).
75
+
76
+ The test is that the op ACCEPTS `sign`. Upstream fla exports a `chunk_kda`
77
+ of the same name that computes the unsigned recurrence; calling it with a
78
+ signed checkpoint's weights would return a plausible tensor rather than an
79
+ error, so the presence of the module is not enough.
80
+ """
81
+ try:
82
+ import inspect
83
+
84
+ from fla.ops.kda import chunk_kda, fused_recurrent_kda
85
+ except Exception:
86
+ return None, None
87
+ try:
88
+ if "sign" not in inspect.signature(chunk_kda).parameters:
89
+ return None, None
90
+ if "sign" not in inspect.signature(fused_recurrent_kda).parameters:
91
+ return None, None
92
+ except (TypeError, ValueError):
93
+ return None, None
94
+ return chunk_kda, fused_recurrent_kda
95
+
96
+
97
+ _CHUNK_KDA, _FUSED_RECURRENT_KDA = _detect_fla()
98
+
99
+
100
+ def _triton_launchable() -> bool:
101
+ """Stricter than importing triton: fla imports fine on a CPU-only box and
102
+ only fails when a kernel is launched."""
103
+ try:
104
+ import triton # noqa: F401
105
+
106
+ return torch.cuda.is_available()
107
+ except Exception:
108
+ return False
109
+
110
+
111
+ HAS_KERNEL = _CHUNK_KDA is not None and _triton_launchable()
112
+
113
+ _REQUESTED = os.environ.get("COMPLEX_KDA_BACKEND", "auto").lower()
114
+ if _REQUESTED not in ("auto", "torch", "kernel"):
115
+ raise ValueError(f"COMPLEX_KDA_BACKEND must be 'auto', 'torch' or 'kernel'; got {_REQUESTED!r}")
116
+ if _REQUESTED == "kernel" and not HAS_KERNEL:
117
+ raise ImportError(
118
+ "COMPLEX_KDA_BACKEND=kernel, but the ComplexKDA fla fork's signed kernels are not "
119
+ "available (need a CUDA device, triton, and `pip install "
120
+ "git+https://github.com/automl/ComplexKDA`).")
121
+ USE_KERNEL = HAS_KERNEL and _REQUESTED != "torch"
122
+
123
+ # Chunk length of the portable recurrence. It trades memory for sequential
124
+ # steps: the intra-chunk term materialises a [chunk, chunk, head_dim] block per
125
+ # head, so 64 is a few tens of MB at these geometries and 256 is a few hundred.
126
+ # It does not change what is computed -- only the order the same sums are taken
127
+ # in, which at fp32 moves a logit by ~1e-6 per layer.
128
+ CHUNK_SIZE = int(os.environ.get("COMPLEX_KDA_CHUNK_SIZE", "64"))
129
+ if CHUNK_SIZE <= 0:
130
+ raise ValueError(f"COMPLEX_KDA_CHUNK_SIZE must be positive; got {CHUNK_SIZE}")
131
+
132
+ if not USE_KERNEL:
133
+ logger.warning_once(
134
+ "ComplexKDA is running its portable torch implementation. For the Triton kernels the "
135
+ "models were trained with, install the fork: "
136
+ "`pip install git+https://github.com/automl/ComplexKDA` (and set "
137
+ "COMPLEX_KDA_BACKEND=torch to keep this path).")
138
+
139
+
140
+ def _fla_modules():
141
+ """fla's fused ShortConvolution / RMSNorm / gated RMSNorm, or (None,)*3."""
142
+ if not USE_KERNEL:
143
+ return None, None, None
144
+ try:
145
+ from fla.modules import FusedRMSNormGated, RMSNorm, ShortConvolution
146
+
147
+ return ShortConvolution, RMSNorm, FusedRMSNormGated
148
+ except Exception:
149
+ return None, None, None
150
+
151
+
152
+ _FLA_SHORTCONV, _FLA_RMSNORM, _FLA_RMSNORM_GATED = _fla_modules()
153
+
154
+
155
+ # ===========================================================================
156
+ # the gate
157
+ #
158
+ # name alpha range activation
159
+ # "softplus" (0, 1] -exp(A_log) * softplus(u)
160
+ # "sigmoid" (0, 1] lower_bound * sigmoid(A * u)
161
+ # "signed_sigmoid2" [-1, 1] 2*sigmoid(u) - 1, evaluated as tanh(u/2)
162
+ # "signed_tanh" [-1, 1] tanh(u)
163
+ #
164
+ # The published baselines use "sigmoid"; the ComplexKDA arms use
165
+ # "signed_sigmoid2".
166
+ # ===========================================================================
167
+
168
+ GATES = ("softplus", "sigmoid", "signed_sigmoid2", "signed_tanh")
169
+
170
+
171
+ def is_signed(gate: str) -> bool:
172
+ if gate not in GATES:
173
+ raise ValueError(f"gate must be one of {GATES}, got {gate!r}")
174
+ return gate.startswith("signed_")
175
+
176
+
177
+ def safe_gate_ok(gate: str) -> bool:
178
+ """Whether log|alpha| is bounded below by `lower_bound`. False only for
179
+ "softplus", which is unbounded."""
180
+ return gate != "softplus"
181
+
182
+
183
+ def signed_gate(z, A_log=None, dt_bias=None, lower_bound: float = -5.0, activation: str = "sigmoid2"):
184
+ """One pre-activation -> (sign, log|alpha|), for alpha in [-1, 1].
185
+
186
+ |alpha| = eps + (1 - eps) * |a|, eps = exp(lower_bound), a = tanh(u/2) or
187
+ tanh(u). "sigmoid2" is `2*sigmoid(u) - 1` written as `tanh(u/2)`: the literal
188
+ spelling cancels catastrophically near u = 0, where the SIGN is decided, so
189
+ it would be settled by rounding rather than by u. The sign shares z with the
190
+ magnitude and is locally constant, so detaching it is exact.
191
+ """
192
+ eps = math.exp(lower_bound)
193
+ u = z.float()
194
+ if dt_bias is not None:
195
+ u = u + dt_bias.view(*([1] * (z.dim() - 2)), *z.shape[-2:])
196
+ if A_log is not None:
197
+ u = A_log.float().exp().view(*([1] * (z.dim() - 2)), -1, 1) * u
198
+ if activation == "sigmoid2":
199
+ a = torch.tanh(0.5 * u)
200
+ elif activation == "tanh":
201
+ a = torch.tanh(u)
202
+ else:
203
+ raise ValueError(f"activation must be 'sigmoid2' or 'tanh', got {activation!r}")
204
+ s = torch.where(a.detach() < 0, -1, 1).to(torch.int8)
205
+ return s, (eps + (1.0 - eps) * a.abs()).log()
206
+
207
+
208
+ def compute_gate(gate, z, A_log=None, dt_bias=None, lower_bound=-5.0):
209
+ """name -> (sign int8 or None, log|alpha| fp32). A None sign is what tells
210
+ the caller there is no gauge to apply."""
211
+ if gate not in GATES:
212
+ raise ValueError(f"gate must be one of {GATES}, got {gate!r}")
213
+ if gate.startswith("signed_"):
214
+ return signed_gate(z, A_log, dt_bias, lower_bound, activation=gate[len("signed_"):])
215
+ u = z.float()
216
+ if dt_bias is not None:
217
+ u = u + dt_bias.view(*([1] * (z.dim() - 2)), *z.shape[-2:])
218
+ A = A_log.float().exp().view(*([1] * (z.dim() - 2)), -1, 1) if A_log is not None else 1.0
219
+ if gate == "softplus":
220
+ return None, -A * F.softplus(u)
221
+ return None, lower_bound * torch.sigmoid(A * u)
222
+
223
+
224
+ def signed_gate_init(dt, lower_bound: float = -5.0, activation: str = "sigmoid2"):
225
+ """dt_bias giving alpha = +exp(-dt) at step 0."""
226
+ eps = math.exp(lower_bound)
227
+ target = ((torch.exp(-dt) - eps) / (1 - eps)).clamp(1e-7, 1 - 1e-7)
228
+ inv = torch.atanh(target)
229
+ return 2.0 * inv if activation == "sigmoid2" else inv
230
+
231
+
232
+ def gate_init(gate, dt, lower_bound=-5.0):
233
+ """dt_bias init inverting each gate's own forward, so all four gates start
234
+ at the same alpha = exp(-dt)."""
235
+ if gate.startswith("signed_"):
236
+ return signed_gate_init(dt, lower_bound, gate[len("signed_"):])
237
+ if gate == "sigmoid":
238
+ p = (dt / abs(lower_bound)).clamp(1e-7, 1 - 1e-7)
239
+ return torch.log(p) - torch.log1p(-p)
240
+ return dt + torch.log(-torch.expm1(-dt))
241
+
242
+
243
+ def init_dt_bias(gate, gate_dim=None, lower_bound=-5.0, gate_init_style="shipped", dt=None):
244
+ if dt is None:
245
+ dt = torch.exp(
246
+ torch.rand(gate_dim, dtype=torch.float32) * (math.log(0.1) - math.log(0.001)) + math.log(0.001)
247
+ ).clamp(min=1e-4)
248
+ init = gate_init(gate, dt, lower_bound)
249
+ if is_signed(gate) and gate_init_style == "spread":
250
+ # Same |alpha| as "shipped" with the sign flipped on half the channels:
251
+ # the activation is odd, so sign and magnitude do not trade off.
252
+ init = init * torch.where(torch.rand_like(init) < 0.5, -1.0, 1.0)
253
+ return init
254
+
255
+
256
+ # ===========================================================================
257
+ # the gauge: carry the +-1 part of alpha as a running sign on q/k
258
+ # ===========================================================================
259
+
260
+ def running_sign(s: torch.Tensor, cu_seqlens: torch.Tensor | None = None) -> torch.Tensor:
261
+ """P_t = prod_{u<=t} s_u along dim 1, as int8.
262
+
263
+ An integer parity prefix sum: exact at any length, and it carries no
264
+ autograd graph, because the sign has no gradient. Resets at sequence starts
265
+ when `cu_seqlens` is given.
266
+ """
267
+ bits = (s < 0).to(torch.int32)
268
+ par = bits.cumsum(dim=1)
269
+ if cu_seqlens is not None:
270
+ starts = cu_seqlens[:-1]
271
+ idx = torch.repeat_interleave(starts, cu_seqlens[1:] - starts)
272
+ par = par - (par[:, idx] - bits[:, idx])
273
+ return torch.where(par & 1 == 1, -1, 1).to(torch.int8)
274
+
275
+
276
+ class _ApplySign(torch.autograd.Function):
277
+ """x * P, keeping P as int8 rather than letting `mul` upcast it."""
278
+
279
+ @staticmethod
280
+ def forward(ctx, x, P):
281
+ ctx.save_for_backward(P)
282
+ return x * P.to(x.dtype)
283
+
284
+ @staticmethod
285
+ def backward(ctx, go):
286
+ (P,) = ctx.saved_tensors
287
+ return go * P.to(go.dtype), None
288
+
289
+
290
+ def apply_sign(x, P):
291
+ return _ApplySign.apply(x, P)
292
+
293
+
294
+ def ungauge_state(ht, P_last, state_v_first: bool, head_k_dim: int | None = None):
295
+ """S_T = Diag(P_T) S~_T, on whichever axis holds K.
296
+
297
+ `state_v_first=True` stores [N, HV, V, K] -- K LAST -- so the axis differs
298
+ between the kernel and the torch path; `head_k_dim` turns a silent
299
+ wrong-axis bug into an assert.
300
+ """
301
+ if ht is None or P_last is None:
302
+ return ht
303
+ if P_last.ndim != ht.ndim - 1:
304
+ raise AssertionError(
305
+ f"gauge rank mismatch: state {tuple(ht.shape)} takes a gauge of "
306
+ f"{ht.ndim - 1} dims, got {tuple(P_last.shape)}.")
307
+ axis = -1 if state_v_first else -2
308
+ if head_k_dim is not None and ht.shape[axis] != head_k_dim:
309
+ raise AssertionError(
310
+ f"state layout mismatch: state_v_first={state_v_first} implies K on axis {axis}, "
311
+ f"but state shape {tuple(ht.shape)} has {ht.shape[axis]} there, not "
312
+ f"head_k_dim={head_k_dim}.")
313
+ P = P_last.to(ht.dtype)
314
+ return ht * (P.unsqueeze(-2) if state_v_first else P.unsqueeze(-1))
315
+
316
+
317
+ # ===========================================================================
318
+ # the recurrence, in torch
319
+ #
320
+ # Both functions take q/k ALREADY l2-normalised and gauged, `g` as log|alpha|,
321
+ # and `beta` already through its sigmoid -- the same contract as the reference
322
+ # implementation in the fork, so the two can be compared term by term.
323
+ # ===========================================================================
324
+
325
+ def recurrent_kda_torch(
326
+ q: torch.Tensor,
327
+ k: torch.Tensor,
328
+ v: torch.Tensor,
329
+ g: torch.Tensor,
330
+ beta: torch.Tensor,
331
+ scale: float | None = None,
332
+ initial_state: torch.Tensor | None = None,
333
+ output_final_state: bool = False,
334
+ ):
335
+ """The definition, one step at a time. [B,T,H,K] q/k, [B,T,HV,V] v,
336
+ [B,T,HV,K] g, [B,T,HV] beta; state [B,HV,K,V]."""
337
+ dtype = v.dtype
338
+ B, T, H, K = q.shape
339
+ HV, V = v.shape[2], v.shape[-1]
340
+ G = HV // H
341
+ if scale is None:
342
+ scale = K ** -0.5
343
+
344
+ q, k, v, g, beta = (x.float() for x in (q, k, v, g, beta))
345
+ q = q.repeat_interleave(G, dim=2) * scale
346
+ k = k.repeat_interleave(G, dim=2)
347
+
348
+ S = q.new_zeros(B, HV, K, V)
349
+ if initial_state is not None:
350
+ S = S + initial_state.float()
351
+ o = torch.zeros_like(v)
352
+ for i in range(T):
353
+ q_i, k_i, v_i, g_i, b_i = q[:, i], k[:, i], v[:, i], g[:, i], beta[:, i]
354
+ S = S * g_i[..., None].exp()
355
+ # delta rule: replace the memory currently read out by k_i with v_i
356
+ S = S + torch.einsum("bhk,bhv->bhkv", b_i[..., None] * k_i, v_i - (k_i[..., None] * S).sum(-2))
357
+ o[:, i] = torch.einsum("bhk,bhkv->bhv", q_i, S)
358
+ return o.to(dtype), (S if output_final_state else None)
359
+
360
+
361
+ def chunk_kda_torch(
362
+ q: torch.Tensor,
363
+ k: torch.Tensor,
364
+ v: torch.Tensor,
365
+ g: torch.Tensor,
366
+ beta: torch.Tensor,
367
+ scale: float | None = None,
368
+ initial_state: torch.Tensor | None = None,
369
+ output_final_state: bool = False,
370
+ chunk_size: int = 64,
371
+ ):
372
+ """The same recurrence in chunks: O(T/C) sequential steps instead of O(T).
373
+
374
+ The WY/UT transform of the chunk's delta updates, then one state carry per
375
+ chunk. Arithmetically identical to `recurrent_kda_torch` up to floating
376
+ point; it exists because a 4096-token forward through the step loop is
377
+ minutes rather than milliseconds.
378
+
379
+ MASK BEFORE EXPONENTIATING. Every exponent used here is a sum of `log|alpha|`
380
+ over an interval, so it is <= 0 and `exp` is safe -- but only for the pairs
381
+ the causal mask keeps. The reference implementation exponentiates the full
382
+ block and masks afterwards, which overflows once `|log alpha| * chunk`
383
+ passes ~88 in fp32: finite forward, NaN backward. Masking first removes that
384
+ failure mode entirely, which is why `chunk_size` needs no upper bound here.
385
+ """
386
+ dtype = v.dtype
387
+ B, T, H, K = q.shape
388
+ HV, V = v.shape[2], v.shape[-1]
389
+ G = HV // H
390
+ if scale is None:
391
+ scale = K ** -0.5
392
+ BT = int(chunk_size)
393
+ if BT <= 0:
394
+ raise ValueError(f"chunk_size must be positive, got {chunk_size}")
395
+
396
+ q, k, v, g, beta = (x.float() for x in (q, k, v, g, beta))
397
+ q = q.repeat_interleave(G, dim=2) * scale
398
+ k = k.repeat_interleave(G, dim=2)
399
+
400
+ # Pad the tail to a whole chunk. beta = 0 makes the padded steps write
401
+ # nothing and g = 0 makes them decay nothing, so the carried state is
402
+ # exactly the state at T.
403
+ pad = (-T) % BT
404
+ if pad:
405
+ q = F.pad(q, (0, 0, 0, 0, 0, pad))
406
+ k = F.pad(k, (0, 0, 0, 0, 0, pad))
407
+ v = F.pad(v, (0, 0, 0, 0, 0, pad))
408
+ g = F.pad(g, (0, 0, 0, 0, 0, pad))
409
+ beta = F.pad(beta, (0, 0, 0, pad))
410
+ NT = (T + pad) // BT
411
+
412
+ # [B, T, HV, X] -> [B, HV, NT, BT, X]
413
+ def _chunks(x):
414
+ return x.view(B, NT, BT, *x.shape[2:]).permute(0, 3, 1, 2, *range(4, x.dim() + 1))
415
+
416
+ q, k, v, g = (_chunks(x) for x in (q, k, v, g))
417
+ beta = beta.view(B, NT, BT, HV).permute(0, 3, 1, 2)
418
+
419
+ eye = torch.eye(BT, device=q.device, dtype=q.dtype)
420
+ rows = torch.arange(BT, device=q.device)
421
+ strictly_lower = rows[:, None] > rows[None, :] # c > i
422
+ causal = rows[:, None] >= rows[None, :] # c >= j
423
+ neg_inf = torch.finfo(q.dtype).min
424
+
425
+ S = q.new_zeros(B, HV, K, V)
426
+ if initial_state is not None:
427
+ S = S + initial_state.float()
428
+ o = torch.zeros_like(v)
429
+
430
+ for n in range(NT):
431
+ q_n, k_n, v_n, g_n, b_n = q[:, :, n], k[:, :, n], v[:, :, n], g[:, :, n], beta[:, :, n]
432
+ gc = g_n.cumsum(-2) # [B,HV,BT,K], <= 0
433
+
434
+ # T[c,i] = beta_c * <k_c * alpha(i,c], k_i> for c > i -- the
435
+ # strictly-lower part of the chunk's own delta interactions.
436
+ d = gc.unsqueeze(-2) - gc.unsqueeze(-3) # [B,HV,BT(c),BT(i),K]
437
+ d = d.masked_fill(~strictly_lower[..., None], neg_inf)
438
+ A = (k_n.unsqueeze(-2) * d.exp() * k_n.unsqueeze(-3)).sum(-1)
439
+ del d
440
+ A = -(A * b_n[..., :, None])
441
+
442
+ # (I - A)^{-1}, A strictly lower and hence unit-triangular after +I.
443
+ # The reference walks the Neumann series row by row; a triangular solve
444
+ # is the same matrix and vectorises.
445
+ Ainv = torch.linalg.solve_triangular(eye - A, eye.expand_as(A), upper=False, unitriangular=True)
446
+ Aw = Ainv * b_n[..., None, :]
447
+
448
+ w = Aw @ (gc.exp() * k_n) # [B,HV,BT,K]
449
+ u = Aw @ v_n # [B,HV,BT,V]
450
+
451
+ dq = gc.unsqueeze(-2) - gc.unsqueeze(-3)
452
+ dq = dq.masked_fill(~causal[..., None], neg_inf)
453
+ Aqk = (q_n.unsqueeze(-2) * dq.exp() * k_n.unsqueeze(-3)).sum(-1)
454
+ del dq
455
+
456
+ v_new = u - w @ S
457
+ o[:, :, n] = (q_n * gc.exp()) @ S + Aqk @ v_new
458
+
459
+ g_last = gc[:, :, -1] # [B,HV,K]
460
+ S = S * g_last.unsqueeze(-1).exp()
461
+ S = S + ((g_last.unsqueeze(-2) - gc).exp() * k_n).transpose(-1, -2) @ v_new
462
+
463
+ o = o.permute(0, 2, 3, 1, 4).reshape(B, NT * BT, HV, V)
464
+ if pad:
465
+ o = o[:, :T]
466
+ return o.to(dtype), (S if output_final_state else None)
467
+
468
+
469
+ # ===========================================================================
470
+ # portable modules
471
+ # ===========================================================================
472
+
473
+ class ShortConvolution(nn.Conv1d):
474
+ """Causal depthwise conv1d with an optional silu.
475
+
476
+ Subclasses nn.Conv1d exactly as the fork's does, so the parameter names
477
+ match and a checkpoint is portable between this path and the Triton one.
478
+ """
479
+
480
+ def __init__(self, hidden_size, kernel_size=4, bias=False, activation="silu"):
481
+ super().__init__(hidden_size, hidden_size, kernel_size, groups=hidden_size, bias=bias)
482
+ if activation not in (None, "silu", "swish"):
483
+ raise ValueError(f"unsupported activation {activation!r}")
484
+ self.hidden_size, self.activation = hidden_size, activation
485
+
486
+ def forward(self, x, cache=None, output_final_state=False, cu_seqlens=None, **kwargs):
487
+ if cu_seqlens is not None:
488
+ raise NotImplementedError(
489
+ "variable-length batching (cu_seqlens) needs the ComplexKDA fla fork")
490
+ B, T, D = x.shape
491
+ w = self.kernel_size[0]
492
+ h = x.transpose(1, 2)
493
+ if cache is not None:
494
+ h = torch.cat([cache, h], dim=-1)[:, :, -(T + w - 1):]
495
+ pad = w - 1 - (h.shape[-1] - T)
496
+ if pad > 0:
497
+ h = F.pad(h, (pad, 0))
498
+ else:
499
+ h = F.pad(h, (w - 1, 0))
500
+ new_cache = h[:, :, -(w - 1):].contiguous() if output_final_state else None
501
+ y = self._conv_forward(h, self.weight, self.bias)[:, :, :T].transpose(1, 2)
502
+ if self.activation in ("silu", "swish"):
503
+ y = F.silu(y)
504
+ return y, new_cache
505
+
506
+
507
+ class RMSNorm(nn.Module):
508
+ """rms(x) * weight, with the fork's optional fused residual add.
509
+
510
+ `forward(x, residual, prenorm=True)` returns `(norm(x + residual), x + residual)`.
511
+ The add is done in the input dtype, matching the fused kernel called with
512
+ `residual_in_fp32=False`.
513
+ """
514
+
515
+ def __init__(self, hidden_size: int, eps: float = 1e-5, elementwise_affine: bool = True):
516
+ super().__init__()
517
+ self.hidden_size, self.eps, self.elementwise_affine = hidden_size, eps, elementwise_affine
518
+ self.weight = nn.Parameter(torch.ones(hidden_size)) if elementwise_affine else None
519
+
520
+ def reset_parameters(self):
521
+ if self.weight is not None:
522
+ nn.init.ones_(self.weight)
523
+
524
+ def extra_repr(self) -> str:
525
+ return f"{self.hidden_size}, eps={self.eps}"
526
+
527
+ def forward(self, x, residual=None, prenorm: bool = False, residual_in_fp32: bool = False):
528
+ if residual is not None:
529
+ x = x + (residual.float() if residual_in_fp32 else residual)
530
+ dt = x.dtype
531
+ xf = x.float()
532
+ y = xf * torch.rsqrt(xf.pow(2).mean(-1, keepdim=True) + self.eps)
533
+ if self.weight is not None:
534
+ y = y * self.weight.float()
535
+ y = y.to(dt)
536
+ return (y, x) if prenorm else y
537
+
538
+
539
+ class FusedRMSNormGated(nn.Module):
540
+ """rms(x) * weight * act(g). The gate is applied AFTER normalising, which is
541
+ what the fused kernel does and is not interchangeable with gating first."""
542
+
543
+ def __init__(self, hidden_size, elementwise_affine=True, eps=1e-5, activation="swish"):
544
+ super().__init__()
545
+ if activation not in ("swish", "silu", "sigmoid"):
546
+ raise ValueError(f"Unsupported activation: {activation}")
547
+ self.hidden_size, self.eps, self.activation = hidden_size, eps, activation
548
+ self.weight = nn.Parameter(torch.ones(hidden_size)) if elementwise_affine else None
549
+
550
+ def reset_parameters(self):
551
+ if self.weight is not None:
552
+ nn.init.ones_(self.weight)
553
+
554
+ def extra_repr(self) -> str:
555
+ return f"{self.hidden_size}, eps={self.eps}, activation={self.activation}"
556
+
557
+ def forward(self, x, g, **kwargs):
558
+ dt = x.dtype
559
+ xf = x.float()
560
+ y = xf * torch.rsqrt(xf.pow(2).mean(-1, keepdim=True) + self.eps)
561
+ if self.weight is not None:
562
+ y = y * self.weight.float()
563
+ gf = g.float()
564
+ y = y * (torch.sigmoid(gf) if self.activation == "sigmoid" else gf * torch.sigmoid(gf))
565
+ return y.to(dt)
566
+
567
+
568
+ class GatedMLP(nn.Module):
569
+ """SwiGLU: down_proj(swish(gate_proj(x)) * up_proj(x))."""
570
+
571
+ def __init__(self, hidden_size: int, hidden_ratio: int | None = None,
572
+ intermediate_size: int | None = None, hidden_act: str = "swish", **kwargs):
573
+ super().__init__()
574
+ if hidden_ratio is None:
575
+ hidden_ratio = 4
576
+ if intermediate_size is None:
577
+ intermediate_size = int(hidden_size * hidden_ratio * 2 / 3)
578
+ intermediate_size = 256 * ((intermediate_size + 256 - 1) // 256)
579
+ if hidden_act not in ("swish", "silu"):
580
+ raise ValueError(f"Unsupported hidden_act: {hidden_act}")
581
+ self.hidden_size, self.hidden_ratio = hidden_size, hidden_ratio
582
+ self.intermediate_size, self.hidden_act = intermediate_size, hidden_act
583
+ self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
584
+ self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
585
+ self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False)
586
+
587
+ def forward(self, x, **kwargs):
588
+ return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
589
+
590
+
591
+ def _rotate_half(x):
592
+ x1, x2 = x.chunk(2, dim=-1)
593
+ return torch.cat((-x2, x1), dim=-1)
594
+
595
+
596
+ class RotaryEmbedding(nn.Module):
597
+ """Rotary position embedding, half-split convention, applied in fp32.
598
+
599
+ Only reached by `use_rope=True` configs; every hybrid published here is
600
+ NoPE, because the linear layers already carry position.
601
+ """
602
+
603
+ def __init__(self, dim: int, base: float = 10000.0):
604
+ super().__init__()
605
+ self.dim, self.base = dim, base
606
+ inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim))
607
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
608
+
609
+ def forward(self, q, k, seqlen_offset=0, max_seqlen=None, cu_seqlens=None):
610
+ if cu_seqlens is not None:
611
+ raise NotImplementedError("variable-length rotary needs the ComplexKDA fla fork")
612
+ T = q.shape[1]
613
+ if torch.is_tensor(seqlen_offset):
614
+ pos = seqlen_offset.view(-1, 1) + torch.arange(T, device=q.device)
615
+ else:
616
+ pos = (torch.arange(T, device=q.device) + int(seqlen_offset)).unsqueeze(0)
617
+ freqs = pos.float().unsqueeze(-1) * self.inv_freq.to(q.device)
618
+ emb = torch.cat((freqs, freqs), dim=-1) # [B or 1, T, dim]
619
+ cos, sin = emb.cos().unsqueeze(-2), emb.sin().unsqueeze(-2)
620
+ qf, kf = q.float(), k.float()
621
+ q = (qf * cos + _rotate_half(qf) * sin).to(q.dtype)
622
+ k = (kf * cos + _rotate_half(kf) * sin).to(k.dtype)
623
+ return q, k
624
+
625
+
626
+ class Attention(nn.Module):
627
+ """The attention layers of a hybrid arm.
628
+
629
+ Causal, through torch SDPA. Two options are not Llama's and both are on in
630
+ the published hybrids: `output_gate` is Qwen3-Next's sigmoid computed from
631
+ the LAYER INPUT and applied before `o_proj` (gating after it would scale the
632
+ residual contribution instead of the per-head mixture, and `o_proj` mixes
633
+ heads, so the two differ), and `use_rope=False` is NoPE -- no rotary at all,
634
+ as in Kimi's hybrid.
635
+ """
636
+
637
+ def __init__(self, hidden_size: int = 2048, num_heads: int = 32, num_kv_heads: int | None = None,
638
+ qkv_bias: bool = False, qk_norm: bool = False, output_gate: bool = False,
639
+ use_rope: bool = True, window_size: int | None = None,
640
+ rope_theta: float | None = 10000.0, max_position_embeddings: int | None = None,
641
+ layer_idx: int | None = None):
642
+ super().__init__()
643
+ self.hidden_size = hidden_size
644
+ self.num_heads = num_heads
645
+ self.num_kv_heads = num_heads if num_kv_heads is None else num_kv_heads
646
+ self.num_kv_groups = num_heads // self.num_kv_heads
647
+ self.head_dim = hidden_size // num_heads
648
+ self.kv_dim = self.num_kv_heads * self.head_dim
649
+ self.qkv_bias, self.qk_norm, self.output_gate, self.use_rope = qkv_bias, qk_norm, output_gate, use_rope
650
+ self.window_size, self.rope_theta = window_size, rope_theta
651
+ self.max_position_embeddings, self.layer_idx = max_position_embeddings, layer_idx
652
+
653
+ self.q_proj = nn.Linear(hidden_size, hidden_size, bias=qkv_bias)
654
+ self.k_proj = nn.Linear(hidden_size, self.kv_dim, bias=qkv_bias)
655
+ self.v_proj = nn.Linear(hidden_size, self.kv_dim, bias=qkv_bias)
656
+ self.o_proj = nn.Linear(hidden_size, hidden_size, bias=False)
657
+ if output_gate:
658
+ self.g_proj = nn.Linear(hidden_size, hidden_size, bias=False)
659
+ if qk_norm:
660
+ self.q_norm = RMSNorm(self.head_dim)
661
+ self.k_norm = RMSNorm(self.head_dim)
662
+ self.rotary = RotaryEmbedding(dim=self.head_dim, base=self.rope_theta) if use_rope else None
663
+
664
+ def forward(self, hidden_states, attention_mask=None, past_key_values=None,
665
+ output_attentions: bool = False, use_cache: bool = False, **kwargs):
666
+ if attention_mask is not None and attention_mask.dim() != 2:
667
+ raise ValueError(
668
+ "Expected attention_mask as a 0-1 matrix of shape [batch_size, seq_len] "
669
+ "(0 = padding). Arbitrary [b, q, k] masks are not supported.")
670
+ if kwargs.get("cu_seqlens") is not None:
671
+ raise NotImplementedError("variable-length attention needs the ComplexKDA fla fork")
672
+
673
+ B, q_len, _ = hidden_states.shape
674
+ q = self.q_proj(hidden_states).view(B, q_len, self.num_heads, self.head_dim)
675
+ k = self.k_proj(hidden_states).view(B, q_len, self.num_kv_heads, self.head_dim)
676
+ v = self.v_proj(hidden_states).view(B, q_len, self.num_kv_heads, self.head_dim)
677
+ if self.qk_norm:
678
+ q, k = self.q_norm(q), self.k_norm(k)
679
+
680
+ seqlen_offset = 0
681
+ if past_key_values is not None:
682
+ seqlen_offset = past_key_values.get_seq_length(self.layer_idx)
683
+ if attention_mask is not None:
684
+ # Padding sits on the LEFT of a padded batch, so a row's real
685
+ # position is its offset minus its padding.
686
+ lens = attention_mask.sum(-1, dtype=torch.long)
687
+ seqlen_offset = seqlen_offset + lens - attention_mask.shape[-1]
688
+ if self.rotary is not None:
689
+ q, k = self.rotary(q, k, seqlen_offset=seqlen_offset)
690
+
691
+ if past_key_values is not None:
692
+ k, v = past_key_values.update_attn(self.layer_idx, k, v, window_size=self.window_size)
693
+
694
+ # [B, T, H, D] -> [B, H, T, D]
695
+ qt, kt, vt = (x.transpose(1, 2) for x in (q, k, v))
696
+ k_len = kt.shape[2]
697
+
698
+ attn_bias = None
699
+ is_causal = False
700
+ if q_len == k_len and attention_mask is None and self.window_size is None:
701
+ is_causal = True
702
+ else:
703
+ pos_q = torch.arange(k_len - q_len, k_len, device=q.device)
704
+ pos_k = torch.arange(k_len, device=q.device)
705
+ # Bottom-right alignment: query t attends keys <= its own position.
706
+ keep = pos_k[None, :] <= pos_q[:, None]
707
+ if self.window_size is not None:
708
+ keep &= pos_k[None, :] > pos_q[:, None] - self.window_size
709
+ keep = keep[None, None]
710
+ if attention_mask is not None:
711
+ pad = attention_mask[:, None, None, :].bool()
712
+ if pad.shape[-1] != k_len:
713
+ pad = F.pad(pad, (k_len - pad.shape[-1], 0), value=True)
714
+ keep = keep & pad
715
+ attn_bias = torch.zeros(keep.shape, dtype=qt.dtype, device=q.device)
716
+ attn_bias = attn_bias.masked_fill(~keep, torch.finfo(qt.dtype).min)
717
+
718
+ gqa = {"enable_gqa": True} if self.num_kv_groups > 1 else {}
719
+ o = F.scaled_dot_product_attention(qt, kt, vt, attn_mask=attn_bias, is_causal=is_causal, **gqa)
720
+ o = o.transpose(1, 2).reshape(B, q_len, -1)
721
+ if self.output_gate:
722
+ o = o * torch.sigmoid(self.g_proj(hidden_states))
723
+ return self.o_proj(o), None, past_key_values
724
+
725
+
726
+ # ===========================================================================
727
+ # the mixer
728
+ # ===========================================================================
729
+
730
+ def _identity(x):
731
+ return x
732
+
733
+
734
+ class ComplexKimiDeltaAttention(nn.Module):
735
+ """Kimi Delta Attention whose decay gate may be negative.
736
+
737
+ Beyond KimiDeltaAttention:
738
+ * `gate` selects the decay parameterisation (two unsigned, two signed);
739
+ * the `+-1` part of a signed gate is carried as a running sign pushed onto
740
+ q/k (the gauge), so the recurrence itself still runs on `|alpha|`. On the
741
+ Triton path the sign is handed to the op as `sign=` and applied inside
742
+ KDA's own l2norm epilogue; the torch path gauges explicitly here.
743
+ """
744
+
745
+ def __init__(self, hidden_size: int = 2048, expand_v: float = 1, head_dim: int = 128,
746
+ num_heads: int = 16, num_v_heads: int | None = None, mode: str = "chunk",
747
+ use_short_conv: bool = True, allow_neg_eigval: bool = False,
748
+ gate: str = "signed_sigmoid2", drop_silu: bool = False, drop_key_silu: bool = False,
749
+ conv_silu: str = "qkv", gate_init_style: str = "shipped",
750
+ output_gate: str = "lowrank", beta_init_style: str = "standard",
751
+ lower_bound: float = -5.0, conv_size: int = 4, conv_bias: bool = False,
752
+ layer_idx: int | None = None, norm_eps: float = 1e-5,
753
+ chunk_size: int | None = None, **kwargs):
754
+ super().__init__()
755
+
756
+ if gate not in GATES:
757
+ raise ValueError(f"gate must be one of {GATES}, got {gate!r}")
758
+ if gate_init_style not in ("shipped", "spread"):
759
+ raise ValueError(f"gate_init_style must be 'shipped' or 'spread', got {gate_init_style!r}")
760
+ if output_gate not in ("lowrank", "linear"):
761
+ raise ValueError(f"output_gate must be 'lowrank' or 'linear', got {output_gate!r}")
762
+ if beta_init_style not in ("standard", "spread"):
763
+ raise ValueError(f"beta_init_style must be 'standard' or 'spread', got {beta_init_style!r}")
764
+ if mode not in ("chunk", "fused_recurrent"):
765
+ raise ValueError(f"unsupported mode {mode!r}")
766
+ if not (-5 <= lower_bound < 0):
767
+ raise ValueError(f"lower_bound must be in [-5, 0), got {lower_bound}")
768
+
769
+ self.mode = mode
770
+ self.allow_neg_eigval = allow_neg_eigval
771
+ self.gate = gate
772
+ self.act = _identity if drop_silu else F.silu
773
+ self.k_act = _identity if drop_silu or drop_key_silu else F.silu
774
+ self.gate_init_style = gate_init_style
775
+ self.beta_init_style = beta_init_style
776
+ self.safe_gate = safe_gate_ok(gate)
777
+ self.lower_bound = lower_bound
778
+ self.hidden_size = hidden_size
779
+ self.expand_v = expand_v
780
+ self.chunk_size = CHUNK_SIZE if chunk_size is None else chunk_size
781
+
782
+ self.use_short_conv = use_short_conv
783
+ self.conv_size = conv_size
784
+ self.conv_bias = conv_bias
785
+ self.head_dim = head_dim
786
+ self.num_heads = num_heads
787
+ self.num_v_heads = num_heads if num_v_heads is None else num_v_heads
788
+
789
+ self.head_k_dim = head_dim
790
+ self.head_v_dim = int(head_dim * expand_v)
791
+ self.key_dim = int(self.num_heads * self.head_k_dim)
792
+ self.value_dim = int(self.num_v_heads * self.head_v_dim)
793
+ self.layer_idx = layer_idx
794
+
795
+ if not math.isclose(head_dim * expand_v, self.head_v_dim, rel_tol=1e-5):
796
+ raise ValueError(f"expand_v={expand_v} does not give an integer head_v_dim from head_dim={head_dim}")
797
+ if self.num_v_heads > self.num_heads and self.num_v_heads % self.num_heads != 0:
798
+ raise ValueError(f"num_v_heads={self.num_v_heads} must be divisible by num_heads={self.num_heads}")
799
+ if self.num_v_heads > self.num_heads and is_signed(gate):
800
+ warnings.warn(
801
+ "signed gate under GVA expands q/k to num_v_heads (the gauge is per value head "
802
+ "but q/k are shared), losing the GVA memory saving.", stacklevel=2)
803
+
804
+ self.q_proj = nn.Linear(hidden_size, self.key_dim, bias=False)
805
+ self.k_proj = nn.Linear(hidden_size, self.key_dim, bias=False)
806
+ self.v_proj = nn.Linear(hidden_size, self.value_dim, bias=False)
807
+
808
+ if any(c not in "qkv" for c in conv_silu):
809
+ raise ValueError(f"conv_silu must be a subset of 'qkv', got {conv_silu!r}")
810
+ self.conv_silu = "" if drop_silu else conv_silu
811
+ if drop_key_silu:
812
+ self.conv_silu = self.conv_silu.replace("k", "")
813
+
814
+ conv_cls = _FLA_SHORTCONV or ShortConvolution
815
+ if use_short_conv:
816
+ self.q_conv1d = conv_cls(hidden_size=self.key_dim, kernel_size=conv_size, bias=conv_bias,
817
+ activation="silu" if "q" in self.conv_silu else None)
818
+ self.k_conv1d = conv_cls(hidden_size=self.key_dim, kernel_size=conv_size, bias=conv_bias,
819
+ activation="silu" if "k" in self.conv_silu else None)
820
+ self.v_conv1d = conv_cls(hidden_size=self.value_dim, kernel_size=conv_size, bias=conv_bias,
821
+ activation="silu" if "v" in self.conv_silu else None)
822
+
823
+ self.gate_dim = int(self.num_v_heads * self.head_k_dim)
824
+ self.f_proj = nn.Sequential(
825
+ nn.Linear(hidden_size, self.head_v_dim, bias=False),
826
+ nn.Linear(self.head_v_dim, self.gate_dim, bias=False),
827
+ )
828
+ self.b_proj = nn.Linear(hidden_size, self.num_v_heads, bias=beta_init_style == "spread")
829
+
830
+ if self.safe_gate:
831
+ self.A_log = nn.Parameter(torch.zeros(self.num_v_heads, dtype=torch.float32))
832
+ else:
833
+ self.A_log = nn.Parameter(torch.log(torch.empty(self.num_v_heads, dtype=torch.float32).uniform_(1, 16)))
834
+ self.A_log._no_weight_decay = True
835
+ self.dt_bias = nn.Parameter(init_dt_bias(gate, self.gate_dim, lower_bound, gate_init_style))
836
+ self.dt_bias._no_weight_decay = True
837
+
838
+ # The output forget gate. "lowrank" is fla's and Kimi Linear's factored
839
+ # hidden -> head_v_dim -> value_dim pair; "linear" is Kimi K3's single
840
+ # full-rank map. Both sit downstream of the recurrence, so neither
841
+ # interacts with the signed decay gate.
842
+ if output_gate == "lowrank":
843
+ self.g_proj = nn.Sequential(
844
+ nn.Linear(hidden_size, self.head_v_dim, bias=False),
845
+ nn.Linear(self.head_v_dim, self.value_dim, bias=True),
846
+ )
847
+ else:
848
+ self.g_proj = nn.Linear(hidden_size, self.value_dim, bias=True)
849
+ norm_gated_cls = _FLA_RMSNORM_GATED or FusedRMSNormGated
850
+ self.o_norm = norm_gated_cls(self.head_v_dim, activation="sigmoid", eps=norm_eps)
851
+ self.o_proj = nn.Linear(self.value_dim, hidden_size, bias=False)
852
+
853
+ def forward(self, hidden_states, attention_mask=None, past_key_values=None,
854
+ use_cache: bool | None = False, output_attentions: bool | None = False, **kwargs):
855
+ if attention_mask is not None and attention_mask.dim() != 2:
856
+ raise ValueError(
857
+ "Expected attention_mask as a 0-1 matrix of shape [batch_size, seq_len] "
858
+ "(0 = padding). Arbitrary [b, q, k] masks are not supported.")
859
+ if kwargs.get("cu_seqlens") is not None and not USE_KERNEL:
860
+ raise NotImplementedError("variable-length batching needs the ComplexKDA fla fork")
861
+ cu_seqlens = kwargs.get("cu_seqlens")
862
+
863
+ B, q_len, _ = hidden_states.shape
864
+ last_state = None
865
+ if past_key_values is not None and self.layer_idx is not None:
866
+ last_state = past_key_values.get(self.layer_idx)
867
+
868
+ if self.use_short_conv:
869
+ cq, ck, cv = last_state["conv_state"] if last_state is not None else (None, None, None)
870
+ q, cq = self.q_conv1d(x=self.q_proj(hidden_states), cache=cq,
871
+ output_final_state=use_cache, cu_seqlens=cu_seqlens)
872
+ k, ck = self.k_conv1d(x=self.k_proj(hidden_states), cache=ck,
873
+ output_final_state=use_cache, cu_seqlens=cu_seqlens)
874
+ v, cv = self.v_conv1d(x=self.v_proj(hidden_states), cache=cv,
875
+ output_final_state=use_cache, cu_seqlens=cu_seqlens)
876
+ else:
877
+ cq = ck = cv = None
878
+ q = self.act(self.q_proj(hidden_states))
879
+ k = self.k_act(self.k_proj(hidden_states))
880
+ v = self.act(self.v_proj(hidden_states))
881
+
882
+ g = self.f_proj(hidden_states)
883
+ beta = self.b_proj(hidden_states)
884
+
885
+ q = q.view(*q.shape[:-1], -1, self.head_k_dim)
886
+ k = k.view(*k.shape[:-1], -1, self.head_k_dim)
887
+ g = g.view(*g.shape[:-1], -1, self.head_k_dim)
888
+ v = v.view(*v.shape[:-1], -1, self.head_v_dim)
889
+
890
+ sign, g = compute_gate(self.gate, g, self.A_log, self.dt_bias, self.lower_bound)
891
+
892
+ # GVA: the gauge is per value head but q/k are shared across the group,
893
+ # so q/k are expanded to HV before either path applies it.
894
+ if sign is not None and sign.shape[2] != q.shape[2]:
895
+ r = sign.shape[2] // q.shape[2]
896
+ q, k = q.repeat_interleave(r, dim=2), k.repeat_interleave(r, dim=2)
897
+
898
+ recurrent_state = last_state["recurrent_state"] if last_state is not None else None
899
+ scale = self.head_k_dim ** -0.5
900
+ P = None
901
+
902
+ if USE_KERNEL:
903
+ # The kernels take `sign` directly: the gauge rides KDA's own l2norm
904
+ # epilogue, and the final state comes back already un-gauged.
905
+ mode = "fused_recurrent" if (q_len <= 64 and not self.training) else self.mode
906
+ op = _CHUNK_KDA if mode == "chunk" else _FUSED_RECURRENT_KDA
907
+ extra = dict(use_gate_in_kernel=False, safe_gate=self.safe_gate) if mode == "chunk" else {}
908
+ o, recurrent_state = op(
909
+ q=q, k=k, v=v, g=g, beta=beta, sign=sign, scale=scale,
910
+ initial_state=recurrent_state, output_final_state=bool(use_cache),
911
+ use_qk_l2norm_in_kernel=True, use_beta_sigmoid_in_kernel=True,
912
+ allow_neg_eigval=self.allow_neg_eigval, lower_bound=self.lower_bound,
913
+ state_v_first=True, cu_seqlens=cu_seqlens, **extra)
914
+ state_v_first = True
915
+ else:
916
+ if sign is not None:
917
+ P = running_sign(sign, cu_seqlens)
918
+ q, k = apply_sign(q, P), apply_sign(k, P)
919
+ qn = F.normalize(q.float(), dim=-1, eps=1e-6).to(q.dtype)
920
+ kn = F.normalize(k.float(), dim=-1, eps=1e-6).to(k.dtype)
921
+ bt = torch.sigmoid(beta.float()) * (2.0 if self.allow_neg_eigval else 1.0)
922
+ fn = recurrent_kda_torch if q_len <= 8 else chunk_kda_torch
923
+ extra = {} if fn is recurrent_kda_torch else dict(chunk_size=min(self.chunk_size, max(q_len, 1)))
924
+ o, recurrent_state = fn(
925
+ qn, kn, v, g.to(q.dtype), bt.to(q.dtype), scale=scale,
926
+ initial_state=recurrent_state, output_final_state=bool(use_cache), **extra)
927
+ state_v_first = False
928
+
929
+ if P is not None and recurrent_state is not None:
930
+ recurrent_state = ungauge_state(recurrent_state, P[:, -1], state_v_first=state_v_first,
931
+ head_k_dim=self.head_k_dim)
932
+
933
+ if use_cache and past_key_values is not None and self.layer_idx is not None:
934
+ past_key_values.update_recurrent(
935
+ self.layer_idx,
936
+ recurrent_state=recurrent_state,
937
+ conv_state=(cq, ck, cv) if self.use_short_conv else None,
938
+ state_v_first=state_v_first,
939
+ offset=q_len,
940
+ )
941
+
942
+ g_out = self.g_proj(hidden_states)
943
+ o = self.o_norm(o, g_out.view(*g_out.shape[:-1], -1, self.head_v_dim))
944
+ o = o.reshape(*o.shape[:-2], -1)
945
+ return self.o_proj(o), None, past_key_values
946
+
947
+
948
+ # ===========================================================================
949
+ # cache
950
+ # ===========================================================================
951
+
952
+ class ComplexKDACache:
953
+ """Per-layer state for incremental decoding.
954
+
955
+ Not a `transformers.Cache`: that class models a growing key/value pair per
956
+ layer, and a linear-attention layer has a FIXED-SIZE recurrent state plus a
957
+ short-convolution window instead. Hybrid arms hold both kinds, which is why
958
+ the two kinds of entry live side by side here.
959
+
960
+ A state is stored in whichever layout produced it (`state_v_first` records
961
+ which), so a cache filled by the Triton path and one filled by the torch
962
+ path are not interchangeable -- the flag makes that a loud error instead of
963
+ a transposed state.
964
+ """
965
+
966
+ # Attributes transformers' generation loop probes on whatever cache it was
967
+ # handed. They are plain class attributes rather than properties so that a
968
+ # version which reads one this does not define fails on the name it wants
969
+ # rather than on something further downstream.
970
+ is_compileable = False
971
+ is_sliding = False
972
+
973
+ def __init__(self, seen_tokens: int = 0):
974
+ self.states: dict[int, dict[str, Any]] = {}
975
+ self._seen_tokens = seen_tokens
976
+
977
+ def __len__(self) -> int:
978
+ return len(self.states)
979
+
980
+ def get(self, layer_idx: int):
981
+ return self.states.get(layer_idx)
982
+
983
+ def get_seq_length(self, layer_idx: int = 0) -> int:
984
+ state = self.states.get(layer_idx)
985
+ return 0 if state is None else state.get("offset", 0)
986
+
987
+ def get_max_cache_shape(self, layer_idx: int = 0) -> int | None:
988
+ return None
989
+
990
+ def update_recurrent(self, layer_idx: int, recurrent_state, conv_state,
991
+ state_v_first: bool, offset: int):
992
+ prev = self.states.get(layer_idx)
993
+ if prev is not None and prev.get("state_v_first") != state_v_first:
994
+ raise ValueError(
995
+ f"layer {layer_idx}: cached state was written with state_v_first="
996
+ f"{prev.get('state_v_first')} and is being updated with {state_v_first}. "
997
+ "The kernel and torch backends store the state on opposite axes; do not "
998
+ "switch COMPLEX_KDA_BACKEND part-way through a generation.")
999
+ self.states[layer_idx] = {
1000
+ "recurrent_state": recurrent_state,
1001
+ "conv_state": conv_state,
1002
+ "state_v_first": state_v_first,
1003
+ "offset": self.get_seq_length(layer_idx) + offset,
1004
+ }
1005
+
1006
+ def update_attn(self, layer_idx: int, k: torch.Tensor, v: torch.Tensor,
1007
+ window_size: int | None = None):
1008
+ """Append these keys/values and return the full history.
1009
+
1010
+ The offset counts TOKENS SEEN, not calls: a prefill hands over many at
1011
+ once, and under a sliding window it keeps counting after the cache has
1012
+ stopped growing. It is what `Attention` rotates by, so getting it from
1013
+ `k.shape[1]` would put a windowed model's rotary back at the start of
1014
+ the window on every step.
1015
+ """
1016
+ prev = self.states.get(layer_idx)
1017
+ n_new = k.shape[1]
1018
+ if prev is not None and prev.get("attn_state") is not None:
1019
+ pk, pv = prev["attn_state"]
1020
+ k, v = torch.cat([pk, k], dim=1), torch.cat([pv, v], dim=1)
1021
+ # TRIM WHAT IS STORED, RETURN THE WHOLE CONCATENATION. Trimming before
1022
+ # the caller attends would hand a prefill of T > window only the last
1023
+ # `window` keys for ALL T queries -- the early ones would then attend a
1024
+ # window that starts after them. The caller applies the window mask;
1025
+ # this only bounds what the NEXT step has to carry, and since what was
1026
+ # stored is already within the window, the concatenation returned on a
1027
+ # decode step is at most `window + 1` long.
1028
+ stored = (k, v) if window_size is None else (k[:, -window_size:], v[:, -window_size:])
1029
+ self.states[layer_idx] = {
1030
+ "attn_state": stored,
1031
+ "offset": (0 if prev is None else prev.get("offset", 0)) + n_new,
1032
+ }
1033
+ return k, v
1034
+
1035
+ def reorder_cache(self, beam_idx: torch.LongTensor):
1036
+ for state in self.states.values():
1037
+ for key in ("recurrent_state",):
1038
+ if state.get(key) is not None:
1039
+ state[key] = state[key].index_select(0, beam_idx.to(state[key].device))
1040
+ if state.get("conv_state") is not None:
1041
+ state["conv_state"] = tuple(
1042
+ None if c is None else c.index_select(0, beam_idx.to(c.device))
1043
+ for c in state["conv_state"])
1044
+ if state.get("attn_state") is not None:
1045
+ state["attn_state"] = tuple(
1046
+ t.index_select(0, beam_idx.to(t.device)) for t in state["attn_state"])
1047
+ return self
1048
+
1049
+
1050
+ # ===========================================================================
1051
+ # the model
1052
+ # ===========================================================================
1053
+
1054
+ class ComplexKDABlock(nn.Module):
1055
+ def __init__(self, config: ComplexKDAConfig, layer_idx: int):
1056
+ super().__init__()
1057
+ self.config = config
1058
+ self.layer_idx = layer_idx
1059
+ norm_cls = _FLA_RMSNORM or RMSNorm
1060
+
1061
+ self.attn_norm = norm_cls(config.hidden_size, eps=config.norm_eps)
1062
+ spec = get_hybrid_attention_spec(config.attn, layer_idx=layer_idx)
1063
+ if spec is not None:
1064
+ # `qk_norm`, `output_gate` and `use_rope` are read with .get: they
1065
+ # are optional keys that the config preserves rather than fields of
1066
+ # the spec, and their defaults are Attention's own. The published
1067
+ # hybrids need the last two -- their attention is GATED and NoPE --
1068
+ # and without them an exported hybrid is a different model: rotary
1069
+ # where the run had none, and no `g_proj` at all.
1070
+ self.attn = Attention(
1071
+ hidden_size=config.hidden_size,
1072
+ num_heads=spec["num_heads"],
1073
+ num_kv_heads=spec["num_kv_heads"],
1074
+ qkv_bias=spec["qkv_bias"],
1075
+ qk_norm=spec.get("qk_norm", False),
1076
+ output_gate=spec.get("output_gate", False),
1077
+ use_rope=spec.get("use_rope", True),
1078
+ window_size=spec["window_size"],
1079
+ rope_theta=spec["rope_theta"],
1080
+ max_position_embeddings=config.max_position_embeddings,
1081
+ layer_idx=layer_idx,
1082
+ )
1083
+ else:
1084
+ self.attn = ComplexKimiDeltaAttention(
1085
+ mode=config.attn_mode,
1086
+ hidden_size=config.hidden_size,
1087
+ expand_v=config.expand_v,
1088
+ head_dim=config.head_dim,
1089
+ num_heads=config.num_heads,
1090
+ num_v_heads=config.num_v_heads,
1091
+ use_short_conv=config.use_short_conv,
1092
+ drop_silu=config.drop_silu,
1093
+ drop_key_silu=config.drop_key_silu,
1094
+ allow_neg_eigval=config.allow_neg_eigval,
1095
+ gate=config.gate,
1096
+ gate_init_style=config.gate_init_style,
1097
+ output_gate=config.output_gate,
1098
+ conv_silu=config.conv_silu,
1099
+ beta_init_style=config.beta_init_style,
1100
+ lower_bound=config.lower_bound,
1101
+ conv_size=config.conv_size,
1102
+ norm_eps=config.norm_eps,
1103
+ layer_idx=layer_idx,
1104
+ )
1105
+ self.mlp_norm = norm_cls(config.hidden_size, eps=config.norm_eps)
1106
+ self.mlp = GatedMLP(
1107
+ hidden_size=config.hidden_size,
1108
+ hidden_ratio=config.hidden_ratio,
1109
+ intermediate_size=config.intermediate_size,
1110
+ hidden_act=config.hidden_act,
1111
+ )
1112
+
1113
+ def forward(self, hidden_states, attention_mask=None, past_key_values=None,
1114
+ use_cache: bool | None = False, output_attentions: bool | None = False, **kwargs):
1115
+ residual = hidden_states
1116
+ hidden_states = self.attn_norm(hidden_states)
1117
+ hidden_states, attentions, past_key_values = self.attn(
1118
+ hidden_states=hidden_states,
1119
+ attention_mask=attention_mask,
1120
+ past_key_values=past_key_values,
1121
+ use_cache=use_cache,
1122
+ output_attentions=output_attentions,
1123
+ **kwargs,
1124
+ )
1125
+ hidden_states, residual = self.mlp_norm(hidden_states, residual, True)
1126
+ hidden_states = self.mlp(hidden_states)
1127
+ return residual + hidden_states, attentions, past_key_values
1128
+
1129
+
1130
+ class ComplexKDAPreTrainedModel(PreTrainedModel):
1131
+ config_class = ComplexKDAConfig
1132
+ base_model_prefix = "model"
1133
+ supports_gradient_checkpointing = True
1134
+ _no_split_modules = ["ComplexKDABlock"]
1135
+ _supports_sdpa = True
1136
+ _can_compile_fullgraph = False
1137
+
1138
+ def _init_weights(self, module: nn.Module):
1139
+ std = self.config.initializer_range
1140
+ if isinstance(module, ComplexKimiDeltaAttention):
1141
+ if next(module.parameters()).device.type != "meta":
1142
+ with torch.no_grad():
1143
+ module.A_log.zero_()
1144
+ dt = torch.exp(
1145
+ torch.rand_like(module.dt_bias) * (math.log(0.1) - math.log(0.001)) + math.log(0.001)
1146
+ ).clamp(min=1e-4)
1147
+ module.dt_bias.copy_(init_dt_bias(
1148
+ module.gate, lower_bound=module.lower_bound,
1149
+ gate_init_style=module.gate_init_style, dt=dt))
1150
+ return
1151
+ if isinstance(module, (nn.Linear, nn.Conv1d)):
1152
+ nn.init.normal_(module.weight, mean=0.0, std=std)
1153
+ if module.bias is not None:
1154
+ nn.init.zeros_(module.bias)
1155
+ elif isinstance(module, nn.Embedding):
1156
+ nn.init.normal_(module.weight, mean=0.0, std=std)
1157
+ elif hasattr(module, "reset_parameters"):
1158
+ module.reset_parameters()
1159
+
1160
+
1161
+ class ComplexKDAModel(ComplexKDAPreTrainedModel):
1162
+ def __init__(self, config: ComplexKDAConfig):
1163
+ super().__init__(config)
1164
+ self.padding_idx = config.pad_token_id
1165
+ self.vocab_size = config.vocab_size
1166
+
1167
+ self.embeddings = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
1168
+ self.layers = nn.ModuleList(
1169
+ [ComplexKDABlock(config, i) for i in range(config.num_hidden_layers)])
1170
+ self.norm = (_FLA_RMSNORM or RMSNorm)(config.hidden_size, eps=config.norm_eps)
1171
+
1172
+ self.gradient_checkpointing = False
1173
+ self.post_init()
1174
+
1175
+ def get_input_embeddings(self):
1176
+ return self.embeddings
1177
+
1178
+ def set_input_embeddings(self, value):
1179
+ self.embeddings = value
1180
+
1181
+ def forward(self, input_ids=None, attention_mask=None, inputs_embeds=None,
1182
+ past_key_values=None, use_cache=None, output_attentions=None,
1183
+ output_hidden_states=None, return_dict=None, **kwargs):
1184
+ if output_attentions:
1185
+ warnings.warn("ComplexKDAModel does not support `output_attentions`; setting it to False.")
1186
+ output_attentions = False
1187
+ output_hidden_states = (output_hidden_states if output_hidden_states is not None
1188
+ else self.config.output_hidden_states)
1189
+ use_cache = use_cache if use_cache is not None else (self.config.use_cache and not self.training)
1190
+ return_dict = return_dict if return_dict is not None else True
1191
+
1192
+ if input_ids is not None and inputs_embeds is not None:
1193
+ raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")
1194
+ if input_ids is None and inputs_embeds is None:
1195
+ raise ValueError("You have to specify either input_ids or inputs_embeds")
1196
+
1197
+ hidden_states = self.embeddings(input_ids) if inputs_embeds is None else inputs_embeds
1198
+
1199
+ if use_cache and past_key_values is None:
1200
+ past_key_values = ComplexKDACache()
1201
+ if past_key_values is not None and not isinstance(past_key_values, ComplexKDACache):
1202
+ raise TypeError(
1203
+ f"ComplexKDA needs a ComplexKDACache (it stores recurrent state, not key/value "
1204
+ f"pairs); got {type(past_key_values).__name__}.")
1205
+
1206
+ all_hidden_states = () if output_hidden_states else None
1207
+ for layer in self.layers:
1208
+ if output_hidden_states:
1209
+ all_hidden_states += (hidden_states,)
1210
+ if self.gradient_checkpointing and self.training:
1211
+ hidden_states, _, past_key_values = self._gradient_checkpointing_func(
1212
+ layer.__call__, hidden_states, attention_mask, past_key_values, use_cache,
1213
+ output_attentions, **kwargs)
1214
+ else:
1215
+ hidden_states, _, past_key_values = layer(
1216
+ hidden_states, attention_mask=attention_mask, past_key_values=past_key_values,
1217
+ use_cache=use_cache, output_attentions=output_attentions, **kwargs)
1218
+
1219
+ hidden_states = self.norm(hidden_states)
1220
+ if output_hidden_states:
1221
+ all_hidden_states += (hidden_states,)
1222
+
1223
+ if not return_dict:
1224
+ return tuple(x for x in (hidden_states, past_key_values, all_hidden_states) if x is not None)
1225
+ return BaseModelOutputWithPast(
1226
+ last_hidden_state=hidden_states,
1227
+ past_key_values=past_key_values,
1228
+ hidden_states=all_hidden_states,
1229
+ attentions=None,
1230
+ )
1231
+
1232
+
1233
+ def _tied_weights_keys_declaration():
1234
+ """How this transformers spells "lm_head.weight IS the embedding".
1235
+
1236
+ Every ladder cell ties its embeddings (the 1.3B arms do not), and the two
1237
+ transformers generations declare that differently:
1238
+
1239
+ 4.x a LIST of regex patterns matched against parameter names
1240
+ 5.x a {target: source} MAPPING
1241
+
1242
+ THE WRONG ONE IS NOT A WARNING. A list under transformers 5 raises
1243
+ `'list' object has no attribute 'keys'` from inside `post_init` -- for
1244
+ every tied checkpoint, and only for tied ones, so it passes every test run
1245
+ against an untied model and then fails for most of the release.
1246
+ """
1247
+ mapping = {"lm_head.weight": "model.embeddings.weight"}
1248
+ patterns = ["lm_head.weight"]
1249
+ try:
1250
+ from transformers.modeling_utils import PreTrainedModel
1251
+
1252
+ annotation = str(getattr(PreTrainedModel, "__annotations__", {})
1253
+ .get("_tied_weights_keys", ""))
1254
+ if annotation:
1255
+ return mapping if "dict" in annotation.lower() else patterns
1256
+ except Exception:
1257
+ pass
1258
+ try:
1259
+ import transformers
1260
+
1261
+ return mapping if int(str(transformers.__version__).split(".")[0]) >= 5 else patterns
1262
+ except Exception:
1263
+ return patterns
1264
+
1265
+
1266
+ class ComplexKDAForCausalLM(ComplexKDAPreTrainedModel, GenerationMixin):
1267
+ _tied_weights_keys = _tied_weights_keys_declaration()
1268
+
1269
+ def __init__(self, config: ComplexKDAConfig):
1270
+ super().__init__(config)
1271
+ self.model = ComplexKDAModel(config)
1272
+ self.vocab_size = config.vocab_size
1273
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
1274
+ self.criterion = None
1275
+ self.post_init()
1276
+
1277
+ def get_input_embeddings(self):
1278
+ return self.model.embeddings
1279
+
1280
+ def set_input_embeddings(self, value):
1281
+ self.model.embeddings = value
1282
+
1283
+ def get_output_embeddings(self):
1284
+ return self.lm_head
1285
+
1286
+ def set_output_embeddings(self, new_embeddings):
1287
+ self.lm_head = new_embeddings
1288
+
1289
+ def get_decoder(self):
1290
+ return self.model
1291
+
1292
+ def set_decoder(self, decoder):
1293
+ self.model = decoder
1294
+
1295
+ def tie_weights(self, *args, **kwargs):
1296
+ """Tie the head to the embedding OURSELVES, rather than describing the
1297
+ tie and hoping this transformers acts on the description.
1298
+
1299
+ Every ladder cell ties, and the exporters write NO `lm_head.weight` for
1300
+ a tied geometry -- there is no second tensor to write. transformers 5.3
1301
+ nonetheless decides the key "is present in the checkpoint", declines to
1302
+ tie, and leaves the head on the META device: `from_pretrained` returns
1303
+ without error and the first `.to(device)` raises "Cannot copy out of
1304
+ meta tensor". A CPU-only smoke test does not even get that far -- it
1305
+ returns a model whose head is data-less.
1306
+
1307
+ Doing the assignment here is version-independent, and transformers
1308
+ calls this both in `post_init` and after loading the weights, so the
1309
+ alias survives materialisation.
1310
+ """
1311
+ if getattr(self.config, "tie_word_embeddings", False):
1312
+ embeddings = self.get_input_embeddings()
1313
+ if embeddings is not None:
1314
+ self.lm_head.weight = embeddings.weight
1315
+ # *args/**kwargs: transformers 5.6 passes `recompute_mapping`, 4.x
1316
+ # passes nothing. Forward whatever it sends rather than pinning a
1317
+ # signature that one of them will not call.
1318
+ return super().tie_weights(*args, **kwargs)
1319
+
1320
+ def forward(self, input_ids=None, attention_mask=None, inputs_embeds=None,
1321
+ past_key_values=None, labels=None, use_cache=None, output_attentions=None,
1322
+ output_hidden_states=None, return_dict=None, logits_to_keep=0, **kwargs):
1323
+ return_dict = return_dict if return_dict is not None else True
1324
+ outputs = self.model(
1325
+ input_ids=input_ids, attention_mask=attention_mask, inputs_embeds=inputs_embeds,
1326
+ past_key_values=past_key_values, use_cache=use_cache,
1327
+ output_attentions=output_attentions, output_hidden_states=output_hidden_states,
1328
+ return_dict=return_dict, **kwargs)
1329
+ hidden_states = outputs[0]
1330
+ if logits_to_keep:
1331
+ hidden_states = hidden_states[:, -logits_to_keep:]
1332
+ logits = self.lm_head(hidden_states)
1333
+
1334
+ loss = None
1335
+ if labels is not None:
1336
+ criterion = self.criterion if self.criterion is not None else nn.CrossEntropyLoss()
1337
+ labels = labels.to(logits.device)
1338
+ # Shift here rather than on the logits, matching the training stack.
1339
+ labels = torch.cat((labels[..., 1:], torch.full_like(labels[:, :1], criterion.ignore_index)), 1)
1340
+ loss = criterion(logits.view(labels.numel(), -1), labels.view(-1))
1341
+
1342
+ if not return_dict:
1343
+ output = (logits,) + tuple(outputs[1:])
1344
+ return (loss,) + output if loss is not None else output
1345
+ return CausalLMOutputWithPast(
1346
+ loss=loss,
1347
+ logits=logits,
1348
+ past_key_values=outputs.past_key_values,
1349
+ hidden_states=outputs.hidden_states,
1350
+ attentions=None,
1351
+ )
1352
+
1353
+ # ---- generation -------------------------------------------------------
1354
+ #
1355
+ # `generate` builds a DynamicCache by default, which this model cannot use
1356
+ # (see ComplexKDACache). Installing ours here is the documented escape
1357
+ # route: a cache already present in model_kwargs is left alone.
1358
+
1359
+ def _prepare_cache_for_generation(self, generation_config, model_kwargs, *args, **kwargs):
1360
+ # *args absorbs the positional tail, which differs across transformers
1361
+ # versions; the two arguments this needs have not moved.
1362
+ if generation_config.use_cache and model_kwargs.get("past_key_values") is None:
1363
+ model_kwargs["past_key_values"] = ComplexKDACache()
1364
+ return True
1365
+ return False
1366
+
1367
+ def prepare_inputs_for_generation(self, input_ids, past_key_values=None, attention_mask=None,
1368
+ inputs_embeds=None, use_cache=True, logits_to_keep=None, **kwargs):
1369
+ if past_key_values is not None and len(past_key_values) > 0:
1370
+ input_ids = input_ids[:, -1:]
1371
+ model_inputs = {"input_ids": input_ids, "inputs_embeds": None}
1372
+ if inputs_embeds is not None and past_key_values is None:
1373
+ model_inputs = {"input_ids": None, "inputs_embeds": inputs_embeds}
1374
+ model_inputs.update(
1375
+ past_key_values=past_key_values,
1376
+ use_cache=use_cache,
1377
+ attention_mask=attention_mask,
1378
+ logits_to_keep=1 if logits_to_keep is None else logits_to_keep,
1379
+ )
1380
+ return model_inputs
1381
+
1382
+ def _reorder_cache(self, past_key_values, beam_idx):
1383
+ return past_key_values.reorder_cache(beam_idx)
special_tokens_map.json ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "bos_token": {
3
+ "content": "</s>",
4
+ "lstrip": false,
5
+ "normalized": true,
6
+ "rstrip": false,
7
+ "single_word": false
8
+ },
9
+ "eos_token": {
10
+ "content": "</s>",
11
+ "lstrip": false,
12
+ "normalized": true,
13
+ "rstrip": false,
14
+ "single_word": false
15
+ },
16
+ "unk_token": {
17
+ "content": "<unk>",
18
+ "lstrip": false,
19
+ "normalized": true,
20
+ "rstrip": false,
21
+ "single_word": false
22
+ }
23
+ }
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer.model ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:9e556afd44213b6bd1be2b850ebbbd98f5481437a8021afaf58ee7fb1818d347
3
+ size 499723
tokenizer_config.json ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_bos_token": false,
3
+ "add_eos_token": false,
4
+ "bos_token": {
5
+ "__type": "AddedToken",
6
+ "content": "</s>",
7
+ "lstrip": false,
8
+ "normalized": true,
9
+ "rstrip": false,
10
+ "single_word": false
11
+ },
12
+ "clean_up_tokenization_spaces": false,
13
+ "chat_template": "{% if messages[0]['role'] == 'system' %}{% set loop_messages = messages[1:] %}{% set system_message = messages[0]['content'] %}{% else %}{% set loop_messages = messages %}{% set system_message = false %}{% endif %}{% for message in loop_messages %}{% if (message['role'] == 'user') != (loop.index0 % 2 == 0) %}{{ raise_exception('Conversation roles must alternate user/assistant/user/assistant/...') }}{% endif %}{% if loop.index0 == 0 and system_message != false %}{% set content = '<<SYS>>\\n' + system_message + '\\n<</SYS>>\\n\\n' + message['content'] %}{% else %}{% set content = message['content'] %}{% endif %}{% if message['role'] == 'user' %}{{ bos_token + '[INST] ' + content.strip() + ' [/INST]' }}{% elif message['role'] == 'assistant' %}{{ ' ' + content.strip() + ' ' + eos_token }}{% endif %}{% endfor %}",
14
+ "eos_token": {
15
+ "__type": "AddedToken",
16
+ "content": "</s>",
17
+ "lstrip": false,
18
+ "normalized": true,
19
+ "rstrip": false,
20
+ "single_word": false
21
+ },
22
+ "model_max_length": 4096,
23
+ "pad_token": null,
24
+ "sp_model_kwargs": {},
25
+ "tokenizer_class": "LlamaTokenizer",
26
+ "unk_token": {
27
+ "__type": "AddedToken",
28
+ "content": "<unk>",
29
+ "lstrip": false,
30
+ "normalized": true,
31
+ "rstrip": false,
32
+ "single_word": false
33
+ }
34
+ }