kylesayrs commited on
Commit
37e4b7a
·
verified ·
1 Parent(s): 0ee605f

Add files using upload-large-folder tool

Browse files
README.md ADDED
@@ -0,0 +1,95 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: mit
3
+ base_model:
4
+ - inference-optimization/Kimi-K3-0.40B
5
+ library_name: compressed-tensors
6
+ tags:
7
+ - quantized
8
+ - mxfp4
9
+ ---
10
+
11
+ # Kimi-K3-0.40B-MXFP4
12
+
13
+ This is an MXFP4-quantized version of [inference-optimization/Kimi-K3-0.40B](https://huggingface.co/inference-optimization/Kimi-K3-0.40B), a tiny model derived from [moonshotai/Kimi-K3](https://huggingface.co/moonshotai/Kimi-K3). Created for testing and development.
14
+
15
+ ## Model Details
16
+
17
+ - **Base Model**: inference-optimization/Kimi-K3-0.40B
18
+ - **Architecture**: kimi_k3
19
+ - **Total Parameters**: 0.40B
20
+ - **Activated Parameters**: ~0.22B (MoE: 2 of 8 experts active per token, plus 1 shared expert)
21
+ - **Quantization**: W4A16 MXFP4 (`mxfp4-pack-quantized`), group size 32
22
+
23
+ ## Quantization Config
24
+
25
+ Matches the quantization scheme used in [moonshotai/Kimi-K3](https://huggingface.co/moonshotai/Kimi-K3/blob/main/config.json):
26
+
27
+ | Field | Value |
28
+ |---|---|
29
+ | Format | `mxfp4-pack-quantized` |
30
+ | Weights | 4-bit float, group_size=32, minmax observer |
31
+ | Scale dtype | `torch.uint8` |
32
+ | Activations | unquantized (W4A16) |
33
+ | Ignored layers | `self_attn`, `shared_experts`, `lm_head`, `vision_tower` |
34
+
35
+ ## Usage
36
+
37
+ ```python
38
+ import torch
39
+ from transformers import AutoModelForCausalLM, AutoTokenizer
40
+ from compressed_tensors.offload import dispatch_model
41
+
42
+ model = AutoModelForCausalLM.from_pretrained(
43
+ "inference-optimization/Kimi-K3-0.40B-MXFP4",
44
+ trust_remote_code=True,
45
+ dtype=torch.bfloat16,
46
+ )
47
+ tokenizer = AutoTokenizer.from_pretrained(
48
+ "inference-optimization/Kimi-K3-0.40B-MXFP4",
49
+ trust_remote_code=True,
50
+ )
51
+ dispatch_model(model)
52
+
53
+ sample = tokenizer("Hello my name is", return_tensors="pt")
54
+ sample = {k: v.to(model.device) for k, v in sample.items()}
55
+ output = model.generate(
56
+ **sample,
57
+ max_new_tokens=100,
58
+ eos_token_id=tokenizer.eos_token_id,
59
+ pad_token_id=tokenizer.pad_token_id,
60
+ )
61
+ print(tokenizer.decode(output[0].tolist(), skip_special_tokens=True))
62
+ ```
63
+
64
+ ## Creation Process
65
+
66
+ Quantized using [llm-compressor](https://github.com/vllm-project/llm-compressor):
67
+
68
+ ```python
69
+ from transformers import AutoModelForCausalLM, AutoTokenizer
70
+ from llmcompressor import oneshot
71
+ from llmcompressor.modifiers.quantization import QuantizationModifier
72
+
73
+ MODEL_ID = "inference-optimization/Kimi-K3-0.40B"
74
+ model = AutoModelForCausalLM.from_pretrained(MODEL_ID, trust_remote_code=True)
75
+ tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, trust_remote_code=True)
76
+
77
+ recipe = QuantizationModifier(
78
+ targets="Linear",
79
+ scheme="MXFP4A16",
80
+ ignore=[
81
+ "re:.*self_attn.*",
82
+ "re:.*shared_experts.*",
83
+ "re:.*lm_head.*",
84
+ "re:.*vision_tower.*",
85
+ ],
86
+ )
87
+ oneshot(model=model, recipe=recipe)
88
+ model.save_pretrained(SAVE_DIR, save_compressed=True)
89
+ tokenizer.save_pretrained(SAVE_DIR)
90
+ ```
91
+
92
+ ## Notes
93
+
94
+ - `trust_remote_code=True` is required to load the custom modeling files.
95
+ - Load with `dtype=torch.bfloat16` to match the decompressed weight dtype.
added_tokens.json ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "<|end_header_id|>": 163844,
3
+ "<|im_assistant|>": 163842,
4
+ "<|im_end|>": 163840,
5
+ "<|im_middle|>": 163846,
6
+ "<|im_system|>": 163845,
7
+ "<|im_user|>": 163841,
8
+ "<|start_header_id|>": 163843
9
+ }
config.json ADDED
@@ -0,0 +1,153 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "KimiK3ForConditionalGeneration"
4
+ ],
5
+ "auto_map": {
6
+ "AutoConfig": "configuration_kimi_k3.KimiK3Config",
7
+ "AutoModel": "modeling_kimi_k3.KimiK3ForConditionalGeneration",
8
+ "AutoModelForCausalLM": "modeling_kimi_k3.KimiK3ForConditionalGeneration"
9
+ },
10
+ "dtype": "float32",
11
+ "ignore_index": -100,
12
+ "media_placeholder_token_id": 163605,
13
+ "model_type": "kimi_k3",
14
+ "pad_token_id": 0,
15
+ "text_config": {
16
+ "_name_or_path": "",
17
+ "activation_situ_beta": 4.0,
18
+ "activation_situ_linear_beta": 25.0,
19
+ "architectures": null,
20
+ "attn_res_block_size": 4,
21
+ "auto_map": {
22
+ "AutoConfig": "configuration_kimi_k3.KimiLinearConfig",
23
+ "AutoModel": "modeling_kimi_linear.KimiLinearModel",
24
+ "AutoModelForCausalLM": "modeling_kimi_linear.KimiLinearForCausalLM"
25
+ },
26
+ "bos_token_id": 1,
27
+ "chunk_size_feed_forward": 0,
28
+ "dtype": null,
29
+ "eos_token_id": 2,
30
+ "first_k_dense_replace": 1,
31
+ "head_dim": 74,
32
+ "hidden_act": "situ",
33
+ "hidden_size": 1024,
34
+ "id2label": {
35
+ "0": "LABEL_0",
36
+ "1": "LABEL_1"
37
+ },
38
+ "initializer_range": 0.02,
39
+ "intermediate_size": 2048,
40
+ "is_encoder_decoder": false,
41
+ "kv_lora_rank": 128,
42
+ "label2id": {
43
+ "LABEL_0": 0,
44
+ "LABEL_1": 1
45
+ },
46
+ "latent_moe_use_norm": true,
47
+ "linear_attn_config": {
48
+ "full_attn_layers": [
49
+ 4,
50
+ 8
51
+ ],
52
+ "head_dim": 32,
53
+ "kda_layers": [
54
+ 1,
55
+ 2,
56
+ 3,
57
+ 5,
58
+ 6,
59
+ 7
60
+ ],
61
+ "num_heads": 8,
62
+ "short_conv_kernel_size": 4,
63
+ "use_full_rank_gate": true
64
+ },
65
+ "max_position_embeddings": 4096,
66
+ "mla_use_nope": true,
67
+ "mla_use_output_gate": true,
68
+ "model_type": "kimi_linear",
69
+ "moe_intermediate_size": 256,
70
+ "moe_layer_freq": 1,
71
+ "moe_renormalize": true,
72
+ "moe_router_activation_func": "sigmoid",
73
+ "num_attention_heads": 8,
74
+ "num_expert_group": 1,
75
+ "num_experts": 8,
76
+ "num_experts_per_token": 2,
77
+ "num_hidden_layers": 8,
78
+ "num_key_value_heads": 8,
79
+ "num_nextn_predict_layers": 0,
80
+ "num_shared_experts": 1,
81
+ "output_attentions": false,
82
+ "output_hidden_states": false,
83
+ "pad_token_id": 0,
84
+ "problem_type": null,
85
+ "q_lora_rank": 256,
86
+ "qk_nope_head_dim": 64,
87
+ "qk_rope_head_dim": 32,
88
+ "return_dict": true,
89
+ "rms_norm_eps": 1e-05,
90
+ "rope_parameters": {
91
+ "rope_theta": 10000.0,
92
+ "rope_type": "default"
93
+ },
94
+ "rope_theta": 10000.0,
95
+ "routed_expert_hidden_size": 512,
96
+ "routed_scaling_factor": 1.0,
97
+ "tie_word_embeddings": false,
98
+ "topk_group": 1,
99
+ "topk_method": "noaux_tc",
100
+ "use_cache": true,
101
+ "use_grouped_topk": true,
102
+ "v_head_dim": 64,
103
+ "vocab_size": 163840
104
+ },
105
+ "transformers_version": "5.15.0.dev0",
106
+ "vision_config": {
107
+ "_name_or_path": "",
108
+ "activation_func": "gelu_pytorch_tanh",
109
+ "architectures": null,
110
+ "attn_bias": false,
111
+ "chunk_size_feed_forward": 0,
112
+ "dtype": null,
113
+ "id2label": {
114
+ "0": "LABEL_0",
115
+ "1": "LABEL_1"
116
+ },
117
+ "init_pos_emb_height": 64,
118
+ "init_pos_emb_time": 4,
119
+ "init_pos_emb_width": 64,
120
+ "is_encoder_decoder": false,
121
+ "label2id": {
122
+ "LABEL_0": 0,
123
+ "LABEL_1": 1
124
+ },
125
+ "linear_bias": false,
126
+ "merge_kernel_size": [
127
+ 2,
128
+ 2
129
+ ],
130
+ "merge_type": "sd2_tpool",
131
+ "mlp_type": "mlp2",
132
+ "mm_hidden_size": 256,
133
+ "mm_projector_type": "patchmergerv2",
134
+ "model_type": "",
135
+ "norm_type": "rmsnorm",
136
+ "output_attentions": false,
137
+ "output_hidden_states": false,
138
+ "patch_embed_proj_bias": false,
139
+ "patch_size": 14,
140
+ "pos_emb_interpolation_mode": "bilinear",
141
+ "pos_emb_type": "divided_fixed",
142
+ "problem_type": null,
143
+ "projector_hidden_act": "gelu",
144
+ "projector_ln_eps": 1e-05,
145
+ "qkv_hidden_size": 1536,
146
+ "return_dict": true,
147
+ "text_hidden_size": 1024,
148
+ "vt_hidden_size": 256,
149
+ "vt_intermediate_size": 512,
150
+ "vt_num_attention_heads": 4,
151
+ "vt_num_hidden_layers": 2
152
+ }
153
+ }
configuration_kimi_k3.py ADDED
@@ -0,0 +1,286 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Optional
2
+
3
+ from transformers.configuration_utils import PretrainedConfig
4
+
5
+
6
+ class KimiLinearConfig(PretrainedConfig):
7
+ model_type = "kimi_linear"
8
+ keys_to_ignore_at_inference = ["past_key_values"]
9
+
10
+ def __init__(
11
+ self,
12
+ model_type="kimi_linear",
13
+ vocab_size=163840,
14
+ hidden_size=4096,
15
+ head_dim=None,
16
+ intermediate_size=11008,
17
+ num_hidden_layers=32,
18
+ num_attention_heads=32,
19
+ num_key_value_heads=None,
20
+ hidden_act="silu",
21
+ initializer_range=0.02,
22
+ rms_norm_eps=1e-6,
23
+ use_cache=True,
24
+ pad_token_id=0,
25
+ bos_token_id=1,
26
+ eos_token_id=2,
27
+ rope_theta=10000.0,
28
+ rope_scaling=None,
29
+ tie_word_embeddings=False,
30
+ moe_intermediate_size: Optional[int] = None,
31
+ moe_renormalize: bool = True,
32
+ moe_router_activation_func: str = "sigmoid",
33
+ num_experts: Optional[int] = None,
34
+ num_experts_per_token: Optional[int] = None,
35
+ num_shared_experts: int = 0,
36
+ routed_scaling_factor: float = 1.0,
37
+ first_k_dense_replace: int = 0,
38
+ moe_layer_freq: int = 1,
39
+ use_grouped_topk: bool = True,
40
+ num_expert_group: int = 1,
41
+ topk_group: int = 1,
42
+ q_lora_rank: Optional[int] = None,
43
+ kv_lora_rank: Optional[int] = None,
44
+ qk_nope_head_dim: Optional[int] = None,
45
+ qk_rope_head_dim: Optional[int] = None,
46
+ v_head_dim: Optional[int] = None,
47
+ mla_use_nope: Optional[bool] = False,
48
+ mla_use_output_gate: Optional[bool] = False,
49
+ num_nextn_predict_layers: int = 0,
50
+ linear_attn_config: Optional[dict] = None,
51
+ attn_res_block_size: Optional[int] = None,
52
+ latent_moe_use_norm: bool = False,
53
+ activation_situ_beta: Optional[float] = None,
54
+ activation_situ_linear_beta: Optional[float] = None,
55
+ max_position_embeddings: int = 4096,
56
+ routed_expert_hidden_size: Optional[int] = None,
57
+ topk_method: str = "noaux_tc",
58
+ **kwargs,
59
+ ):
60
+ self.model_type = model_type
61
+ self.vocab_size = vocab_size
62
+ self.hidden_size = hidden_size
63
+ self.head_dim = (
64
+ head_dim if head_dim is not None else hidden_size // num_attention_heads
65
+ )
66
+ self.intermediate_size = intermediate_size
67
+ self.num_hidden_layers = num_hidden_layers
68
+ self.num_attention_heads = num_attention_heads
69
+
70
+ # for backward compatibility
71
+ if num_key_value_heads is None:
72
+ num_key_value_heads = num_attention_heads
73
+
74
+ self.num_key_value_heads = num_key_value_heads
75
+ self.hidden_act = hidden_act
76
+ self.initializer_range = initializer_range
77
+ self.rms_norm_eps = rms_norm_eps
78
+ self.use_cache = use_cache
79
+ self.rope_theta = rope_theta
80
+ self.rope_scaling = rope_scaling
81
+
82
+ self.q_lora_rank = q_lora_rank
83
+ self.kv_lora_rank = kv_lora_rank
84
+ self.qk_nope_head_dim = qk_nope_head_dim
85
+ self.qk_rope_head_dim = qk_rope_head_dim
86
+ self.v_head_dim = v_head_dim
87
+ self.mla_use_nope = mla_use_nope
88
+ self.mla_use_output_gate = mla_use_output_gate
89
+ # moe config
90
+ self.num_experts = num_experts
91
+ self.num_experts_per_token = num_experts_per_token
92
+ self.moe_renormalize = moe_renormalize
93
+ self.num_shared_experts = num_shared_experts
94
+ self.routed_scaling_factor = routed_scaling_factor
95
+ self.moe_router_activation_func = moe_router_activation_func
96
+ assert self.moe_router_activation_func in ("softmax", "sigmoid")
97
+ self.moe_intermediate_size = moe_intermediate_size
98
+ self.first_k_dense_replace = first_k_dense_replace
99
+ self.moe_layer_freq = moe_layer_freq
100
+ self.use_grouped_topk = use_grouped_topk
101
+ self.num_expert_group = num_expert_group
102
+ self.topk_group = topk_group
103
+ self.num_nextn_predict_layers = num_nextn_predict_layers
104
+
105
+ self.attn_res_block_size = attn_res_block_size
106
+ self.latent_moe_use_norm = latent_moe_use_norm
107
+ self.activation_situ_beta = activation_situ_beta
108
+ self.activation_situ_linear_beta = activation_situ_linear_beta
109
+ self.max_position_embeddings = max_position_embeddings
110
+ self.routed_expert_hidden_size = routed_expert_hidden_size
111
+ self.topk_method = topk_method
112
+
113
+ if linear_attn_config is not None:
114
+ assert linear_attn_config["kda_layers"] is not None
115
+ assert linear_attn_config["full_attn_layers"] is not None
116
+ self.linear_attn_config = linear_attn_config
117
+
118
+ super().__init__(
119
+ pad_token_id=pad_token_id,
120
+ bos_token_id=bos_token_id,
121
+ eos_token_id=eos_token_id,
122
+ tie_word_embeddings=tie_word_embeddings,
123
+ **kwargs,
124
+ )
125
+
126
+ @property
127
+ def is_mla(self):
128
+ return (
129
+ self.q_lora_rank is not None
130
+ or self.kv_lora_rank is not None
131
+ or self.qk_nope_head_dim is not None
132
+ or self.qk_rope_head_dim is not None
133
+ or self.v_head_dim is not None
134
+ or self.mla_use_nope is True
135
+ )
136
+
137
+ @property
138
+ def is_moe(self):
139
+ return self.num_experts is not None
140
+
141
+ @property
142
+ def is_linear_attn(self) -> bool:
143
+ return not (
144
+ self.linear_attn_config is None
145
+ or (
146
+ isinstance(self.linear_attn_config, dict)
147
+ and self.linear_attn_config["kda_layers"] is not None
148
+ and len(self.linear_attn_config["kda_layers"]) == 0
149
+ )
150
+ )
151
+
152
+ def is_kda_layer(self, layer_idx: int):
153
+ return (
154
+ self.linear_attn_config is not None
155
+ and (layer_idx + 1) in self.linear_attn_config["kda_layers"]
156
+ )
157
+
158
+
159
+ class KimiK3VisionConfig(PretrainedConfig):
160
+ def __init__(
161
+ self,
162
+ patch_size: int = 14,
163
+ init_pos_emb_height: int = 64,
164
+ init_pos_emb_width: int = 64,
165
+ init_pos_emb_time: int = 4,
166
+ pos_emb_type: str = "divided_fixed",
167
+ vt_num_attention_heads: int = 12,
168
+ vt_num_hidden_layers: int = 27,
169
+ vt_hidden_size: int = 1024,
170
+ vt_intermediate_size: int = 4096,
171
+ merge_kernel_size: tuple = (2, 2),
172
+ merge_type: str = "sd2_tpool",
173
+ _attn_implementation: str = "flash_attention_2",
174
+ # MM Projector parameters
175
+ mm_projector_type: str = "patchmergerv2",
176
+ mm_hidden_size: int | None = None,
177
+ projector_hidden_act: str = "gelu",
178
+ projector_ln_eps: float = 1e-5,
179
+ # vision tower parameters
180
+ qkv_hidden_size: int = 1536,
181
+ norm_type: str = "rmsnorm",
182
+ attn_bias: bool = False,
183
+ patch_embed_proj_bias: bool = False,
184
+ mlp_type: str = "mlp2",
185
+ linear_bias: bool = False,
186
+ activation_func: str = "gelu_pytorch_tanh",
187
+ pos_emb_interpolation_mode: str = "bilinear",
188
+ # Other parameters
189
+ ignore_index: int = -100,
190
+ media_placeholder_token_id: int = 163605,
191
+ pad_token_id: int = 0,
192
+ text_hidden_size=7168,
193
+ **kwargs,
194
+ ):
195
+ self.patch_size = patch_size
196
+ self.init_pos_emb_height = init_pos_emb_height
197
+ self.init_pos_emb_width = init_pos_emb_width
198
+ self.init_pos_emb_time = init_pos_emb_time
199
+ self.pos_emb_type = pos_emb_type
200
+ self.vt_num_attention_heads = vt_num_attention_heads
201
+ self.vt_num_hidden_layers = vt_num_hidden_layers
202
+ self.vt_hidden_size = vt_hidden_size
203
+ self.vt_intermediate_size = vt_intermediate_size
204
+ self.merge_kernel_size = merge_kernel_size
205
+ self.merge_type = merge_type
206
+ self._attn_implementation = _attn_implementation
207
+
208
+ # MM Projector config
209
+ self.mm_projector_type = mm_projector_type
210
+ self.mm_hidden_size = (
211
+ mm_hidden_size if mm_hidden_size is not None else vt_hidden_size
212
+ )
213
+ self.projector_hidden_act = projector_hidden_act
214
+ self.projector_ln_eps = projector_ln_eps
215
+ self.text_hidden_size = text_hidden_size
216
+
217
+ # vision tower parameters
218
+ self.qkv_hidden_size = qkv_hidden_size
219
+ self.norm_type = norm_type
220
+ self.attn_bias = attn_bias
221
+ self.patch_embed_proj_bias = patch_embed_proj_bias
222
+ self.mlp_type = mlp_type
223
+ self.linear_bias = linear_bias
224
+ self.activation_func = activation_func
225
+ self.pos_emb_interpolation_mode = pos_emb_interpolation_mode
226
+
227
+ super().__init__(**kwargs)
228
+
229
+
230
+ class KimiK3Config(PretrainedConfig):
231
+ """Kimi-K3 model configuration.
232
+
233
+ Args:
234
+ text_config (dict | KimiLinearConfig): Configuration for the text model.
235
+
236
+ Vision Tower Parameters (from MoonViT3dConfig):
237
+ patch_size (int): Patch size for vision tower.
238
+ init_pos_emb_height (int): Initial position embedding height.
239
+ init_pos_emb_width (int): Initial position embedding width.
240
+ init_pos_emb_time (int): Initial position embedding time dimension.
241
+ pos_emb_type (str): Type of position embedding.
242
+ vt_num_attention_heads (int): Number of attention heads in vision tower.
243
+ vt_num_hidden_layers (int): Number of hidden layers in vision tower.
244
+ vt_hidden_size (int): Hidden size of vision tower.
245
+ vt_intermediate_size (int): Intermediate size in vision tower FFN.
246
+ merge_kernel_size (tuple): Kernel size for patch merging.
247
+ merge_type (str): Type of merge operation.
248
+ _attn_implementation (str): Attention implementation type.
249
+
250
+ MM Projector Parameters (from MultiModalProjectorConfig):
251
+ mm_projector_type (str): Type of multimodal projector.
252
+ mm_hidden_size (int): Hidden size from vision tower (should match vt_hidden_size).
253
+ projector_hidden_act (str): Activation function for projector.
254
+ projector_ln_eps (float): Layer norm epsilon for projector.
255
+
256
+ Other Parameters:
257
+ ignore_index (int): The ignore index for the loss function.
258
+ media_placeholder_token_id (int): The token ID to use for media placeholders.
259
+ pad_token_id (int): The token ID to use for padding.
260
+ """
261
+
262
+ model_type = "kimi_k3"
263
+
264
+ def __init__(
265
+ self,
266
+ text_config: dict | KimiLinearConfig = None,
267
+ vision_config: dict | KimiK3VisionConfig = None,
268
+ # Other parameters
269
+ ignore_index: int = -100,
270
+ media_placeholder_token_id: int = 163605,
271
+ pad_token_id: int = 0,
272
+ **kwargs,
273
+ ):
274
+ if isinstance(text_config, dict):
275
+ text_config = KimiLinearConfig(**text_config)
276
+ if isinstance(vision_config, dict):
277
+ vision_config = KimiK3VisionConfig(**vision_config)
278
+ self.text_config = text_config
279
+ self.vision_config = vision_config
280
+ # Other config
281
+ self.ignore_index = ignore_index
282
+ self.media_placeholder_token_id = media_placeholder_token_id
283
+ if getattr(self.text_config, "quantization_config", None) is not None:
284
+ self.quantization_config = self.text_config.quantization_config
285
+
286
+ super().__init__(pad_token_id=pad_token_id, **kwargs)
encoding_k3.py ADDED
@@ -0,0 +1,651 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Kimi K3 XTML encoding helpers.
2
+
3
+ This module keeps chat rendering in Python.
4
+ Callers that need token IDs should consume ``EncodeSegment`` objects directly:
5
+ structural markers may be encoded as tiktoken special tokens, while user/tool
6
+ text and attribute values are encoded as ordinary text.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import json
12
+ from dataclasses import dataclass
13
+ from typing import Any, Iterable, Optional
14
+
15
+ OPEN_TOKEN = "<|open|>"
16
+ CLOSE_TOKEN = "<|close|>"
17
+ SEP_TOKEN = "<|sep|>"
18
+ END_OF_MSG_TOKEN = "<|end_of_msg|>"
19
+ IMAGE_PLACEHOLDER = "<|kimi_image_placeholder|>"
20
+
21
+ _VALID_THINKING_EFFORTS = {"low", "high", "max"}
22
+
23
+
24
+ @dataclass(frozen=True)
25
+ class EncodeSegment:
26
+ text: str
27
+ allow_special: bool = False
28
+
29
+
30
+ class _ImagePromptState:
31
+ def __init__(self, image_prompts: Optional[list[str]] = None):
32
+ self.image_prompts = image_prompts
33
+ self.index = 0
34
+
35
+ def next_prompt(self) -> str:
36
+ if self.image_prompts is None:
37
+ return IMAGE_PLACEHOLDER
38
+ if self.index >= len(self.image_prompts):
39
+ raise ValueError("More image placeholders than image prompts.")
40
+ prompt = self.image_prompts[self.index]
41
+ self.index += 1
42
+ return prompt
43
+
44
+ def assert_consumed(self) -> None:
45
+ if self.image_prompts is None:
46
+ return
47
+ if self.index != len(self.image_prompts):
48
+ raise ValueError(
49
+ f"image prompt count {len(self.image_prompts)} != "
50
+ f"consumed placeholder count {self.index}"
51
+ )
52
+
53
+
54
+ def _segment(text: Any, *, allow_special: bool = False) -> list[EncodeSegment]:
55
+ text = str(text)
56
+ if not text:
57
+ return []
58
+ return [EncodeSegment(text, allow_special=allow_special)]
59
+
60
+
61
+ def _control(text: str) -> list[EncodeSegment]:
62
+ return _segment(text, allow_special=True)
63
+
64
+
65
+ def _text(text: Any) -> list[EncodeSegment]:
66
+ return _segment(text, allow_special=False)
67
+
68
+
69
+ def _append_text(
70
+ segments: list[EncodeSegment],
71
+ text: Any,
72
+ image_state: _ImagePromptState,
73
+ ) -> None:
74
+ text = str(text)
75
+ if text == "":
76
+ return
77
+ if image_state.image_prompts is None or IMAGE_PLACEHOLDER not in text:
78
+ segments.extend(_text(text))
79
+ return
80
+
81
+ parts = text.split(IMAGE_PLACEHOLDER)
82
+ for i, part in enumerate(parts):
83
+ segments.extend(_text(part))
84
+ if i < len(parts) - 1:
85
+ segments.extend(_segment(image_state.next_prompt(), allow_special=True))
86
+
87
+
88
+ def _escape_attr_value(value: Any) -> str:
89
+ return str(value).replace("&", "&amp;").replace('"', "&quot;")
90
+
91
+
92
+ def _attr(key: str, value: Any) -> list[EncodeSegment]:
93
+ return (
94
+ _text(f" {key}") + _text('="') + _text(_escape_attr_value(value)) + _text('"')
95
+ )
96
+
97
+
98
+ def _open_tag(tag: str, attrs: Iterable[tuple[str, Any]] = ()) -> list[EncodeSegment]:
99
+ segments: list[EncodeSegment] = []
100
+ segments.extend(_control(OPEN_TOKEN))
101
+ segments.extend(_text(tag))
102
+ for key, value in attrs:
103
+ segments.extend(_attr(key, value))
104
+ segments.extend(_control(SEP_TOKEN))
105
+ return segments
106
+
107
+
108
+ def _close_tag(tag: str) -> list[EncodeSegment]:
109
+ segments: list[EncodeSegment] = []
110
+ segments.extend(_control(CLOSE_TOKEN))
111
+ segments.extend(_text(tag))
112
+ segments.extend(_control(SEP_TOKEN))
113
+ return segments
114
+
115
+
116
+ def _end_of_msg() -> list[EncodeSegment]:
117
+ return _control(END_OF_MSG_TOKEN)
118
+
119
+
120
+ def _json_compact(value: Any) -> str:
121
+ return json.dumps(value, ensure_ascii=False, separators=(",", ":"))
122
+
123
+
124
+ def _is_mapping(value: Any) -> bool:
125
+ return isinstance(value, dict)
126
+
127
+
128
+ def _xtml_type(value: Any) -> str:
129
+ if isinstance(value, bool):
130
+ return "boolean"
131
+ if value is None:
132
+ return "null"
133
+ if isinstance(value, (int, float)) and not isinstance(value, bool):
134
+ return "number"
135
+ if isinstance(value, str):
136
+ return "string"
137
+ if _is_mapping(value):
138
+ return "object"
139
+ return "array"
140
+
141
+
142
+ def _xtml_value(value: Any) -> str:
143
+ if isinstance(value, str):
144
+ return value
145
+ return json.dumps(value, ensure_ascii=False)
146
+
147
+
148
+ def _get_value(obj: Any, key: str, default: Any = None) -> Any:
149
+ if isinstance(obj, dict):
150
+ return obj.get(key, default)
151
+ return getattr(obj, key, default)
152
+
153
+
154
+ def extract_response_schema(response_format: Any) -> Any:
155
+ if response_format is None:
156
+ return None
157
+
158
+ json_schema = _get_value(response_format, "json_schema")
159
+ if json_schema is None:
160
+ return None
161
+
162
+ if isinstance(json_schema, dict):
163
+ return json_schema.get(
164
+ "schema",
165
+ json_schema.get("json_schema", json_schema),
166
+ )
167
+
168
+ schema = _get_value(json_schema, "schema")
169
+ if schema is not None:
170
+ return schema
171
+
172
+ schema = _get_value(json_schema, "json_schema")
173
+ if schema is not None:
174
+ return schema
175
+
176
+ return json_schema
177
+
178
+
179
+ def deep_sort_dict(obj: Any) -> Any:
180
+ if isinstance(obj, dict):
181
+ return {k: deep_sort_dict(v) for k, v in sorted(obj.items())}
182
+ if isinstance(obj, list):
183
+ return [deep_sort_dict(item) for item in obj]
184
+ return obj
185
+
186
+
187
+ def normalize_tool_arguments(arguments: Any) -> tuple[dict[str, Any], Optional[str]]:
188
+ if arguments is None:
189
+ return {}, None
190
+ if isinstance(arguments, dict):
191
+ return arguments, None
192
+ if isinstance(arguments, str):
193
+ if not arguments.strip():
194
+ return {}, None
195
+ try:
196
+ parsed = json.loads(arguments)
197
+ except json.JSONDecodeError:
198
+ return {}, arguments
199
+ if not isinstance(parsed, dict):
200
+ raise ValueError("Kimi K3 tool call arguments must be a JSON object.")
201
+ return parsed, None
202
+ raise TypeError(
203
+ "Kimi K3 tool call arguments must be a dict or a JSON object string."
204
+ )
205
+
206
+
207
+ def normalize_message(message: Any) -> Any:
208
+ if not isinstance(message, dict):
209
+ return message
210
+
211
+ normalized = dict(message)
212
+
213
+ tools = normalized.get("tools")
214
+ if tools is not None:
215
+ normalized["tools"] = deep_sort_dict(tools)
216
+
217
+ tool_calls = normalized.get("tool_calls")
218
+ if not tool_calls:
219
+ return normalized
220
+
221
+ normalized_calls = []
222
+ for tool_call in tool_calls:
223
+ if not isinstance(tool_call, dict):
224
+ normalized_calls.append(tool_call)
225
+ continue
226
+
227
+ tc = dict(tool_call)
228
+ function = tc.get("function")
229
+ if isinstance(function, dict):
230
+ fn = dict(function)
231
+ arguments, json_block = normalize_tool_arguments(fn.get("arguments"))
232
+ fn["arguments"] = arguments
233
+ if json_block is None:
234
+ fn.pop("_xtml_json_block", None)
235
+ else:
236
+ fn["_xtml_json_block"] = json_block
237
+ tc["function"] = fn
238
+ else:
239
+ arguments, json_block = normalize_tool_arguments(tc.get("arguments"))
240
+ tc["arguments"] = arguments
241
+ if json_block is None:
242
+ tc.pop("_xtml_json_block", None)
243
+ else:
244
+ tc["_xtml_json_block"] = json_block
245
+ normalized_calls.append(tc)
246
+
247
+ normalized["tool_calls"] = normalized_calls
248
+ return normalized
249
+
250
+
251
+ def normalize_conversation(conversation: Any) -> Any:
252
+ if not isinstance(conversation, list):
253
+ return conversation
254
+
255
+ def normalize_messages(messages: list[Any]) -> list[Any]:
256
+ return [normalize_message(message) for message in messages]
257
+
258
+ if conversation and isinstance(conversation[0], list):
259
+ return [normalize_messages(messages) for messages in conversation]
260
+ return normalize_messages(conversation)
261
+
262
+
263
+ def _tool_call_id_index(tool_calls: Any) -> dict:
264
+ """Map assistant ``tool_calls[].id`` to ``(1-based position, function name)``.
265
+
266
+ The position mirrors the chat template's enumeration over ``tool_calls``
267
+ (every entry advances the position, even an id-less one). Duplicate ids keep
268
+ their first occurrence.
269
+ """
270
+ index: dict = {}
271
+ if not isinstance(tool_calls, list):
272
+ return index
273
+ for position, tool_call in enumerate(tool_calls, start=1):
274
+ if not isinstance(tool_call, dict):
275
+ continue
276
+ call_id = tool_call.get("id")
277
+ if call_id is None:
278
+ continue
279
+ key = str(call_id)
280
+ if key in index:
281
+ continue
282
+ function = tool_call.get("function")
283
+ name = (
284
+ function.get("name")
285
+ if isinstance(function, dict)
286
+ else tool_call.get("name")
287
+ )
288
+ index[key] = (position, name)
289
+ return index
290
+
291
+
292
+ def normalize_xtml_tool_result_messages(messages: list[Any]) -> list[Any]:
293
+ """Re-sort K3 XTML tool results into assistant ``tool_calls`` order.
294
+
295
+ Serving frameworks generally deliver tool results already in call order. A
296
+ direct Transformers caller, however, may pass OpenAI-style tool messages in any
297
+ order, so each run of consecutive tool messages is matched against the most
298
+ recent preceding assistant ``tool_calls`` by opaque ``tool_call_id`` ==
299
+ ``tool_calls[].id`` (K3 drops the ``func:index`` format requirement) and
300
+ sorted by the matched 1-based position. The matched call is authoritative,
301
+ so each matched message's ``tool`` is set to that call's function name --
302
+ this keeps an explicit (and possibly stale) ``tool``/``name`` from drifting
303
+ out of sync with the reordered position. ``index`` is still derived from the
304
+ rendered position by the chat template. A run that cannot be fully matched is
305
+ left untouched. Re-running is idempotent.
306
+
307
+ This function is side-effect free: matched tool messages are shallow-copied
308
+ before their ``tool``/``name`` is rewritten, and every other message is
309
+ appended to the output as-is. The input list and its message objects are
310
+ never mutated.
311
+ """
312
+ if not isinstance(messages, list):
313
+ return messages
314
+
315
+ output: list[Any] = []
316
+ current_index: dict = {}
317
+ i = 0
318
+ n = len(messages)
319
+
320
+ while i < n:
321
+ message = messages[i]
322
+
323
+ if isinstance(message, dict) and message.get("role") == "assistant":
324
+ tool_calls = message.get("tool_calls")
325
+ current_index = _tool_call_id_index(tool_calls) if tool_calls else {}
326
+ output.append(message)
327
+ i += 1
328
+ continue
329
+
330
+ if not isinstance(message, dict) or message.get("role") != "tool":
331
+ output.append(message)
332
+ i += 1
333
+ continue
334
+
335
+ run: list[tuple] = [] # (position, original_offset, message, name)
336
+ unresolved = False
337
+ offset = 0
338
+ while (
339
+ i < n
340
+ and isinstance(messages[i], dict)
341
+ and messages[i].get("role") == "tool"
342
+ ):
343
+ tool_message = messages[i]
344
+ call_id = tool_message.get("tool_call_id", tool_message.get("id"))
345
+ matched = current_index.get(str(call_id)) if call_id is not None else None
346
+ if matched is None:
347
+ unresolved = True
348
+ run.append((None, offset, tool_message, None))
349
+ else:
350
+ position, name = matched
351
+ run.append((position, offset, tool_message, name))
352
+ offset += 1
353
+ i += 1
354
+
355
+ if unresolved:
356
+ output.extend(item[2] for item in run)
357
+ else:
358
+ run.sort(key=lambda item: (item[0], item[1]))
359
+ for _, _, tool_message, name in run:
360
+ if name is None:
361
+ output.append(tool_message)
362
+ continue
363
+ # The id-matched call is authoritative: align tool (and any
364
+ # explicit name) so the rendered XTML tool attribute cannot
365
+ # disagree with the reordered position. Copy first so the
366
+ # caller's message object is never mutated.
367
+ resolved = dict(tool_message)
368
+ resolved["tool"] = name
369
+ if "name" in resolved:
370
+ resolved["name"] = name
371
+ output.append(resolved)
372
+
373
+ return output
374
+
375
+
376
+ def is_batched_conversation(conversation: Any) -> bool:
377
+ return (
378
+ isinstance(conversation, list)
379
+ and bool(conversation)
380
+ and isinstance(conversation[0], list)
381
+ )
382
+
383
+
384
+ def _render_content_segments(
385
+ content: Any,
386
+ image_state: _ImagePromptState,
387
+ ) -> list[EncodeSegment]:
388
+ segments: list[EncodeSegment] = []
389
+ if isinstance(content, str):
390
+ _append_text(segments, content, image_state)
391
+ elif content is not None:
392
+ for part in content:
393
+ if part["type"] in ["image", "image_url"]:
394
+ segments.extend(_segment(image_state.next_prompt(), allow_special=True))
395
+ else:
396
+ _append_text(segments, part["text"], image_state)
397
+ return segments
398
+
399
+
400
+ def _internal_system_message(message_type: str, body: str) -> list[EncodeSegment]:
401
+ segments: list[EncodeSegment] = []
402
+ segments.extend(_open_tag("message", [("role", "system"), ("type", message_type)]))
403
+ segments.extend(_text(body.strip()))
404
+ segments.extend(_close_tag("message"))
405
+ segments.extend(_end_of_msg())
406
+ return segments
407
+
408
+
409
+ def _render_assistant_segments(
410
+ message: dict[str, Any],
411
+ image_state: _ImagePromptState,
412
+ thinking: bool = True,
413
+ ) -> list[EncodeSegment]:
414
+ segments: list[EncodeSegment] = []
415
+ # The <think> channel is structural: in thinking mode every assistant
416
+ # message carries the open/close tags even when there is no reasoning
417
+ # content to fill in. In non-thinking mode the channel is dropped
418
+ # entirely.
419
+ if thinking:
420
+ reasoning_content = message.get("reasoning_content") or message.get("reasoning")
421
+ segments.extend(_open_tag("think"))
422
+ if reasoning_content is not None and str(reasoning_content).strip():
423
+ _append_text(segments, reasoning_content, image_state)
424
+ segments.extend(_close_tag("think"))
425
+
426
+ segments.extend(_open_tag("response"))
427
+ segments.extend(_render_content_segments(message.get("content"), image_state))
428
+ segments.extend(_close_tag("response"))
429
+
430
+ tool_calls = message.get("tool_calls")
431
+ if tool_calls:
432
+ segments.extend(_open_tag("tools"))
433
+ for index, tool_call in enumerate(tool_calls, start=1):
434
+ fn = tool_call.get("function", tool_call)
435
+ segments.extend(_open_tag("call", [("tool", fn["name"]), ("index", index)]))
436
+ args = fn.get("arguments", {})
437
+ json_block = fn.get("_xtml_json_block")
438
+ if json_block is not None:
439
+ segments.extend(_open_tag("json", [("type", "object")]))
440
+ _append_text(segments, json_block, image_state)
441
+ segments.extend(_close_tag("json"))
442
+ elif _is_mapping(args):
443
+ for key, value in args.items():
444
+ segments.extend(
445
+ _open_tag(
446
+ "argument",
447
+ [("key", key), ("type", _xtml_type(value))],
448
+ )
449
+ )
450
+ _append_text(segments, _xtml_value(value), image_state)
451
+ segments.extend(_close_tag("argument"))
452
+ segments.extend(_close_tag("call"))
453
+ segments.extend(_close_tag("tools"))
454
+
455
+ return segments
456
+
457
+
458
+ def _render_tool_declare(tools: Any, *, dynamic: bool = False) -> list[EncodeSegment]:
459
+ if dynamic:
460
+ body = (
461
+ "## New Tools Available\n"
462
+ "The system dynamically extends the toolset via lazy-loading.\n"
463
+ "You have access to all existing and extended tools.\n"
464
+ "Here are the specs for the extended tools.\n\n"
465
+ "```json\n"
466
+ f"{_json_compact(tools)}\n"
467
+ "```"
468
+ )
469
+ else:
470
+ body = (
471
+ "# Tools\n"
472
+ "Here are the available tools, described in JSONSchema.\n\n"
473
+ "```json\n"
474
+ f"{_json_compact(tools)}\n"
475
+ "```"
476
+ )
477
+ segments: list[EncodeSegment] = []
478
+ segments.extend(
479
+ _open_tag("message", [("role", "system"), ("type", "tool-declare")])
480
+ )
481
+ segments.extend(_text(body))
482
+ segments.extend(_close_tag("message"))
483
+ segments.extend(_end_of_msg())
484
+ return segments
485
+
486
+
487
+ def build_chat_segments(
488
+ messages: list[Any],
489
+ tools: Optional[list[dict]] = None,
490
+ *,
491
+ add_generation_prompt: bool = True,
492
+ thinking: bool = True,
493
+ image_prompts: Optional[list[str]] = None,
494
+ **kwargs: Any,
495
+ ) -> list[EncodeSegment]:
496
+ # Re-sort tool results by tool_call_id at the lowest layer so every caller
497
+ # (processor or direct tokenizer) gets correctly ordered XTML. The helper is
498
+ # side-effect free, so the caller's message objects are left untouched.
499
+ messages = normalize_xtml_tool_result_messages(messages)
500
+ messages = normalize_conversation(messages)
501
+ tools = deep_sort_dict(tools)
502
+
503
+ kwargs = dict(kwargs)
504
+ response_format = kwargs.get("response_format")
505
+ if "response_schema" not in kwargs:
506
+ response_schema = extract_response_schema(response_format)
507
+ if response_schema is not None:
508
+ kwargs["response_schema"] = response_schema
509
+ if kwargs.get("response_schema") is not None:
510
+ kwargs["response_schema"] = deep_sort_dict(kwargs["response_schema"])
511
+
512
+ image_state = _ImagePromptState(image_prompts)
513
+ segments: list[EncodeSegment] = []
514
+
515
+ tool_calls = None
516
+ tool_index = 0
517
+
518
+ if tools:
519
+ segments.extend(_render_tool_declare(tools))
520
+
521
+ thinking_effort = kwargs.get("thinking_effort")
522
+ if thinking and thinking_effort is not None:
523
+ assert thinking_effort in _VALID_THINKING_EFFORTS, (
524
+ f"Unsupported thinking_effort={thinking_effort!r}; "
525
+ f"supported values are {sorted(_VALID_THINKING_EFFORTS)}."
526
+ )
527
+ if thinking and thinking_effort in _VALID_THINKING_EFFORTS:
528
+ segments.extend(
529
+ _internal_system_message(
530
+ "thinking-effort",
531
+ "`thinking_effort` guides on how much to think in your "
532
+ "thinking channel (not including the response channel), "
533
+ "supported values include `low`, `medium`, `high`, and `max`.\n"
534
+ f"Now the system is invoked with `thinking_effort={thinking_effort}`.",
535
+ )
536
+ )
537
+
538
+ for message_index, message in enumerate(messages):
539
+ if not isinstance(message, dict):
540
+ continue
541
+
542
+ role = message["role"]
543
+ if role == "user":
544
+ attrs = [("role", "user")]
545
+ if message.get("name"):
546
+ attrs.append(("name", message["name"]))
547
+ segments.extend(_open_tag("message", attrs))
548
+ segments.extend(
549
+ _render_content_segments(message.get("content"), image_state)
550
+ )
551
+ segments.extend(_close_tag("message"))
552
+ segments.extend(_end_of_msg())
553
+ elif role == "system" and message.get("tools"):
554
+ segments.extend(_render_tool_declare(message["tools"], dynamic=True))
555
+ elif role == "system":
556
+ attrs = [("role", "system")]
557
+ if message.get("name"):
558
+ attrs.append(("name", message["name"]))
559
+ segments.extend(_open_tag("message", attrs))
560
+ segments.extend(
561
+ _render_content_segments(message.get("content"), image_state)
562
+ )
563
+ segments.extend(_close_tag("message"))
564
+ segments.extend(_end_of_msg())
565
+ elif role == "tool":
566
+ tool_index += 1
567
+ tool_name = message.get("tool", message.get("name"))
568
+ if (
569
+ tool_name is None
570
+ and tool_calls is not None
571
+ and tool_index <= len(tool_calls)
572
+ ):
573
+ tc = tool_calls[tool_index - 1]
574
+ fn = tc.get("function", tc)
575
+ tool_name = fn["name"]
576
+ if tool_name is None:
577
+ raise ValueError(
578
+ "Kimi K3 tool messages need a resolvable tool name: "
579
+ "carry `tool`/`name`, or match a preceding assistant "
580
+ "tool_call by order."
581
+ )
582
+ segments.extend(
583
+ _open_tag(
584
+ "message",
585
+ [("role", "tool"), ("tool", tool_name), ("index", tool_index)],
586
+ )
587
+ )
588
+ segments.extend(
589
+ _render_content_segments(message.get("content"), image_state)
590
+ )
591
+ segments.extend(_close_tag("message"))
592
+ segments.extend(_end_of_msg())
593
+ elif role == "assistant":
594
+ tool_calls = message.get("tool_calls")
595
+ tool_index = 0
596
+ attrs = [("role", "assistant")]
597
+ if message.get("name"):
598
+ attrs.append(("name", message["name"]))
599
+ segments.extend(_open_tag("message", attrs))
600
+ segments.extend(_render_assistant_segments(message, image_state, thinking))
601
+ segments.extend(_close_tag("message"))
602
+ segments.extend(_end_of_msg())
603
+
604
+ tool_choice = kwargs.get("tool_choice")
605
+ if tool_choice == "required":
606
+ segments.extend(
607
+ _internal_system_message(
608
+ "tool-choice",
609
+ "The system is invoked with `tool_choice=required`.\n"
610
+ "You MUST call tools in the next message.",
611
+ )
612
+ )
613
+ elif tool_choice == "none":
614
+ segments.extend(
615
+ _internal_system_message(
616
+ "tool-choice",
617
+ "The system is invoked with `tool_choice=none`.\n"
618
+ "You MUST NOT call any tools in the next message.",
619
+ )
620
+ )
621
+
622
+ rf = kwargs.get("response_format")
623
+ rf_type = _get_value(rf, "type", rf) if isinstance(rf, dict) else rf
624
+ if rf_type == "json_object":
625
+ segments.extend(
626
+ _internal_system_message(
627
+ "response-format",
628
+ "The system is invoked with `response_format=json_object`.\n"
629
+ "Your response must be raw JSON data without markdown code "
630
+ "blocks (```json) or any additional formatting.",
631
+ )
632
+ )
633
+ elif rf_type == "json_schema":
634
+ schema = _json_compact(kwargs.get("response_schema"))
635
+ segments.extend(
636
+ _internal_system_message(
637
+ "response-format",
638
+ "The system is invoked with `response_format=json_schema`.\n"
639
+ "Your response must be raw JSON data without markdown code "
640
+ "blocks (```json) or any additional formatting.\n"
641
+ "The JSON data must match the following schema:\n"
642
+ f"```json\n{schema}\n```",
643
+ )
644
+ )
645
+
646
+ if add_generation_prompt:
647
+ segments.extend(_open_tag("message", [("role", "assistant")]))
648
+ segments.extend(_open_tag("think" if thinking else "response"))
649
+
650
+ image_state.assert_consumed()
651
+ return segments
generation_config.json ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_from_model_config": true,
3
+ "bos_token_id": 1,
4
+ "eos_token_id": 2,
5
+ "output_attentions": false,
6
+ "output_hidden_states": false,
7
+ "pad_token_id": 0,
8
+ "transformers_version": "5.15.0.dev0",
9
+ "use_cache": false
10
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a39b7ed2769ee2f9891a19cef6f1ca7986a3f295c09759b5b3135285e5d7a678
3
+ size 1582323896
modeling_kimi_k3.py ADDED
@@ -0,0 +1,1355 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # coding=utf-8
2
+ # Copyright 2025-2026 The Moonshot AI Team and HuggingFace Inc. team. All rights reserved.
3
+ #
4
+ # The code is based on llava (llava/modeling_llava.py), but modified for Kimi-K3.
5
+ #
6
+ # Licensing Information:
7
+ # - Code derived from llava (llava/modeling_llava.py) is licensed under the Apache License, Version 2.0.
8
+ # - Other parts of the code are licensed under the Kimi K3 License (see the LICENSE file in this repository).
9
+ #
10
+ # Apache License, Version 2.0:
11
+ # Licensed under the Apache License, Version 2.0 (the "License");
12
+ # you may not use this file except in compliance with the License.
13
+ # You may obtain a copy of the License at
14
+ #
15
+ # http://www.apache.org/licenses/LICENSE-2.0
16
+ #
17
+ # Unless required by applicable law or agreed to in writing, software
18
+ # distributed under the License is distributed on an "AS IS" BASIS,
19
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
20
+ # See the License for the specific language governing permissions and
21
+ # limitations under the License.
22
+
23
+
24
+ # NOTE: Reference implementation for model architecture; see the model card for production deployment.
25
+ import math
26
+ from collections.abc import Sequence
27
+ from copy import deepcopy
28
+ from typing import Optional
29
+
30
+ import numpy as np
31
+ import torch
32
+ import torch.nn as nn
33
+ import torch.nn.functional as F
34
+ from transformers import activations
35
+
36
+ try:
37
+ from transformers.activations import PytorchGELUTanh
38
+ except ImportError:
39
+ from transformers.activations import GELUTanh
40
+
41
+ activations.PytorchGELUTanh = GELUTanh
42
+ PytorchGELUTanh = GELUTanh
43
+ from transformers.activations import PytorchGELUTanh
44
+ from transformers.configuration_utils import PretrainedConfig
45
+ from transformers.modeling_utils import PreTrainedModel
46
+ from transformers.models.llava.modeling_llava import LlavaCausalLMOutputWithPast
47
+ from transformers.utils import is_flash_attn_2_available
48
+
49
+ from .configuration_kimi_k3 import KimiK3Config
50
+ from .modeling_kimi_linear import KimiLinearForCausalLM
51
+
52
+ # Flash attention imports
53
+ if is_flash_attn_2_available():
54
+ from flash_attn import flash_attn_varlen_func
55
+ else:
56
+ flash_attn_varlen_func = None
57
+
58
+
59
+ def multihead_attention(
60
+ q: torch.Tensor,
61
+ k: torch.Tensor,
62
+ v: torch.Tensor,
63
+ q_cu_seqlens: torch.Tensor | None = None,
64
+ k_cu_seqlens: torch.Tensor | None = None,
65
+ max_seqlen_q: int | None = None,
66
+ max_seqlen_k: int | None = None,
67
+ deterministic: bool = False,
68
+ ):
69
+ """Multi-head attention using flash attention 2.
70
+
71
+ Args:
72
+ q, k, v: tensor of shape (batch_size, seqlen, num_heads, head_dim),
73
+ or (tot_seqlens, num_heads, head_dim) if packing.
74
+ q_cu_seqlens (torch.Tensor): cumulative sequence lengths of q.
75
+ The first element should be 0 and the last element should be q.shape[0].
76
+ k_cu_seqlens (torch.Tensor): cumulative sequence lengths of k.
77
+ The first element should be 0 and the last element should be k.shape[0].
78
+
79
+ Returns:
80
+ output: shape (batch_size, seqlen, dim) or (tot_seqlens, dim) if packing,
81
+ where dim = num_heads * head_dim
82
+ """
83
+ attn_out = flash_attn_varlen_func(
84
+ q,
85
+ k,
86
+ v,
87
+ q_cu_seqlens,
88
+ k_cu_seqlens,
89
+ max_seqlen_q,
90
+ max_seqlen_k,
91
+ causal=False,
92
+ deterministic=deterministic,
93
+ )
94
+ if isinstance(attn_out, tuple):
95
+ attn_out = attn_out[0]
96
+
97
+ attn_out = attn_out.flatten(start_dim=-2)
98
+
99
+ return attn_out
100
+
101
+
102
+ def eager_attention(
103
+ q: torch.Tensor,
104
+ k: torch.Tensor,
105
+ v: torch.Tensor,
106
+ q_cu_seqlens: Optional[torch.Tensor] = None,
107
+ k_cu_seqlens: Optional[torch.Tensor] = None,
108
+ **kwargs,
109
+ ) -> torch.Tensor:
110
+ seq_length = q.shape[0]
111
+ attention_mask = torch.zeros(
112
+ [1, seq_length, seq_length], device=q.device, dtype=torch.bool
113
+ )
114
+ for i in range(1, len(q_cu_seqlens)):
115
+ attention_mask[
116
+ ...,
117
+ q_cu_seqlens[i - 1] : q_cu_seqlens[i],
118
+ q_cu_seqlens[i - 1] : q_cu_seqlens[i],
119
+ ] = True
120
+ q = q.transpose(0, 1)
121
+ k = k.transpose(0, 1)
122
+ v = v.transpose(0, 1)
123
+
124
+ attn_weight = q @ k.transpose(-2, -1) / math.sqrt(q.shape[-1])
125
+ attn_weight = attn_weight.masked_fill(
126
+ ~attention_mask, torch.finfo(attn_weight.dtype).min
127
+ )
128
+ attn_weight = torch.softmax(attn_weight, dim=-1, dtype=torch.float32).to(q.dtype)
129
+
130
+ attn_output = attn_weight @ v
131
+ attn_output = attn_output.transpose(0, 1)
132
+ attn_output = attn_output.reshape(seq_length, -1)
133
+ return attn_output
134
+
135
+
136
+ VL_VISION_ATTENTION_FUNCTIONS = {
137
+ "flash_attention_2": multihead_attention,
138
+ "eager": eager_attention,
139
+ }
140
+
141
+
142
+ def _apply_rope_input_validation(x, freqs_cis):
143
+ assert x.ndim == freqs_cis.ndim + 1, (x.shape, freqs_cis.shape)
144
+ assert x.shape[:-2] == freqs_cis.shape[:-1], (x.shape, freqs_cis.shape)
145
+ assert x.shape[-1] == 2 * freqs_cis.shape[-1], (x.shape, freqs_cis.shape)
146
+ assert freqs_cis.dtype == torch.complex64, freqs_cis.dtype
147
+
148
+
149
+ def get_rope_shape_decorate(func):
150
+ _get_rope_shape_first_call_flag = set()
151
+
152
+ def wrapper(org, interpolation_mode, shape):
153
+ key = (org.requires_grad, torch.is_grad_enabled(), interpolation_mode)
154
+ if key not in _get_rope_shape_first_call_flag:
155
+ _get_rope_shape_first_call_flag.add(key)
156
+ _ = func(org, interpolation_mode, shape=(64, 64))
157
+ return func(org, interpolation_mode, shape)
158
+
159
+ return wrapper
160
+
161
+
162
+ @get_rope_shape_decorate
163
+ @torch.compile(dynamic=True)
164
+ def get_rope_shape(org, interpolation_mode, shape):
165
+ return (
166
+ F.interpolate(
167
+ org.permute((2, 0, 1)).unsqueeze(0),
168
+ size=shape,
169
+ mode=interpolation_mode,
170
+ )
171
+ .squeeze(0)
172
+ .permute((1, 2, 0))
173
+ .flatten(end_dim=1)
174
+ )
175
+
176
+
177
+ def apply_rope(
178
+ xq: torch.Tensor, xk: torch.Tensor, freqs_cis: torch.Tensor
179
+ ) -> tuple[torch.Tensor, torch.Tensor]:
180
+ """
181
+ Args: (The leading dimensions of all inputs should be the same)
182
+ xq: query, tensor of shape (..., num_heads, head_dim)
183
+ xk: key, tensor of shape (..., num_heads, head_dim)
184
+ freqs_cis: tensor of shape (..., head_dim/2), dtype=torch.complex64. It contains the precomputed cis(freqs) for each position in the 2D grid.
185
+ Returns:
186
+ xq_out, xk_out: tensors of shape (..., num_heads, head_dim)
187
+ """
188
+ _apply_rope_input_validation(xq, freqs_cis)
189
+ _apply_rope_input_validation(xk, freqs_cis)
190
+
191
+ freqs_cis = freqs_cis.unsqueeze(-2) # ..., 1, head_dim/2
192
+ # ..., num_heads, head_dim/2
193
+ xq_ = torch.view_as_complex(xq.float().view(*xq.shape[:-1], -1, 2))
194
+ xk_ = torch.view_as_complex(xk.float().view(*xq.shape[:-1], -1, 2))
195
+ xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(-2) # ..., num_heads, head_dim
196
+ xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(-2) # ..., num_heads, head_dim
197
+ return xq_out.type_as(xq), xk_out.type_as(xk)
198
+
199
+
200
+ def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
201
+ """
202
+ From:
203
+ https://github.com/OpenGVLab/InternVideo/blob/421f6d2361fc8f61a3394244571f2601a4e99e29/InternVideo2/multi_modality/models/backbones/internvideo2/pos_embed.py#L86
204
+ embed_dim: output dimension for each position
205
+ pos: a list of positions to be encoded: size (M,)
206
+ out: (M, D)
207
+ """
208
+ assert embed_dim % 2 == 0
209
+ omega = np.arange(embed_dim // 2, dtype=np.float32)
210
+ omega /= embed_dim / 2.0
211
+ omega = 1.0 / 10000**omega # (D/2,)
212
+
213
+ pos = pos.reshape(-1) # (M,)
214
+ out = np.einsum("m,d->md", pos, omega) # (M, D/2), outer product
215
+
216
+ emb_sin = np.sin(out) # (M, D/2)
217
+ emb_cos = np.cos(out) # (M, D/2)
218
+
219
+ emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
220
+ return emb
221
+
222
+
223
+ def get_1d_sincos_pos_embed(embed_dim, t_size, cls_token=False):
224
+ """
225
+ t_size: int of the temporal size
226
+ return:
227
+ pos_embed: [t_size, embed_dim] or [1+t_size, embed_dim] (w/ or w/o cls_token)
228
+ """
229
+ grid_t = np.arange(t_size, dtype=np.float32)
230
+ pos_embed = get_1d_sincos_pos_embed_from_grid(embed_dim, grid_t)
231
+ if cls_token:
232
+ pos_embed = np.concatenate([np.zeros([1, embed_dim]), pos_embed], axis=0)
233
+ return pos_embed
234
+
235
+
236
+ class Learnable2DInterpPosEmbDivided_fixed(nn.Module):
237
+ def __init__(
238
+ self,
239
+ height: int,
240
+ width: int,
241
+ num_frames: int,
242
+ dim: int,
243
+ interpolation_mode: str = "bicubic",
244
+ ) -> None:
245
+ super().__init__()
246
+ self.height = height
247
+ self.width = width
248
+ self.num_frames = num_frames
249
+ self.dim = dim
250
+ self.interpolation_mode = interpolation_mode
251
+ self.weight = nn.Parameter(torch.empty(height, width, dim))
252
+ self.register_buffer(
253
+ "time_weight",
254
+ torch.from_numpy(get_1d_sincos_pos_embed(self.dim, self.num_frames))
255
+ .float()
256
+ .unsqueeze(1),
257
+ persistent=False,
258
+ )
259
+
260
+ self.reset_parameters()
261
+
262
+ def reset_parameters(self):
263
+ nn.init.normal_(self.weight)
264
+
265
+ def forward(self, x: torch.Tensor, grid_thws: torch.Tensor) -> torch.Tensor:
266
+ pos_embs = []
267
+ for t, h, w in grid_thws.tolist():
268
+ assert t <= self.num_frames, f"t:{t} > self.num_frames:{self.num_frames}"
269
+ if (h, w) == self.weight.shape[:-1]:
270
+ pos_emb_2d = self.weight.flatten(end_dim=1)
271
+ else:
272
+ pos_emb_2d = get_rope_shape(
273
+ self.weight,
274
+ interpolation_mode=self.interpolation_mode,
275
+ shape=(h, w),
276
+ )
277
+
278
+ if t == 1:
279
+ pos_emb_3d = pos_emb_2d
280
+ else:
281
+ pos_emb_3d = (
282
+ pos_emb_2d.unsqueeze(0).repeat(t, 1, 1) + self.time_weight[0:t]
283
+ )
284
+
285
+ pos_embs.append(pos_emb_3d.reshape(-1, pos_emb_3d.shape[-1]))
286
+
287
+ out = x + torch.cat(pos_embs)
288
+ return out
289
+
290
+
291
+ class MoonVision3dPatchEmbed(nn.Module):
292
+ def __init__(
293
+ self,
294
+ out_dim: int,
295
+ in_dim: int = 3,
296
+ patch_size: int | tuple[int, int] = (14, 14),
297
+ pos_emb_height: int = 14,
298
+ pos_emb_width: int = 14,
299
+ pos_emb_time: int = 4,
300
+ pos_emb_type: str = "divided_fixed",
301
+ patch_embed_proj_bias: bool = True,
302
+ pos_emb_interpolation_mode: str = "bicubic",
303
+ ):
304
+ super().__init__()
305
+ assert isinstance(
306
+ patch_size, int | Sequence
307
+ ), f"Invalid patch_size type: {type(patch_size)}"
308
+ if isinstance(patch_size, int):
309
+ patch_size = (patch_size, patch_size)
310
+ assert (
311
+ len(patch_size) == 2
312
+ ), f"Expected patch_size to be a tuple of 2, got {patch_size}"
313
+ self.patch_size = patch_size
314
+
315
+ self.proj = nn.Conv2d(
316
+ in_dim,
317
+ out_dim,
318
+ kernel_size=patch_size,
319
+ stride=patch_size,
320
+ bias=patch_embed_proj_bias,
321
+ )
322
+
323
+ if pos_emb_type == "divided_fixed":
324
+ self.pos_emb = Learnable2DInterpPosEmbDivided_fixed(
325
+ height=pos_emb_height,
326
+ width=pos_emb_width,
327
+ num_frames=pos_emb_time,
328
+ dim=out_dim,
329
+ interpolation_mode=pos_emb_interpolation_mode,
330
+ )
331
+ else:
332
+ raise NotImplementedError(f"Not support pos_emb_type: {pos_emb_type}")
333
+
334
+ def forward(self, x: torch.Tensor, grid_thws: torch.Tensor) -> torch.Tensor:
335
+ """
336
+ Args:
337
+ x (L, Channels): input tensor
338
+ grid_hws (N, 3): temporal, height and width
339
+
340
+ Returns:
341
+ (L, Cout) tensor
342
+ """
343
+ x = self.proj(x).view(x.size(0), -1)
344
+ # apply positional embedding
345
+ x = self.pos_emb(x, grid_thws)
346
+ return x
347
+
348
+
349
+ class Rope2DPosEmbRepeated(nn.Module):
350
+ """2D rotary position embedding with multi-resolution support.
351
+
352
+ This class is intended to be used in the following way:
353
+ 1. Before training, create an instance of Rope2DPosEmb. This instance will hold the precomputed cis.
354
+ 2. Before each forward pass, call `get_freqs_cis_by_*` to get the `freqs_cis` tensor for this iteration.
355
+ 3. During the forward pass, pass the `freqs_cis` tensor to each attention layer, and call `apply` just before each attention operation.
356
+ The rope is shared across all attention layers and all heads.
357
+
358
+ Refs:
359
+ - RoFormer: https://arxiv.org/abs/2104.09864
360
+ - VisionLLaMA: https://arxiv.org/abs/2403.00522
361
+ - https://github.com/Meituan-AutoML/VisionLLaMA/blob/main/dit/models.py
362
+
363
+ Args:
364
+ dim (int): usually the multi-head attention dimension, should be divisible by 4 (TODO: relax this constraint if needed)
365
+ max_height (int): the maximum height of the 2D grid
366
+ max_width (int): the maximum width of the 2D grid
367
+ theta_base (float): the base of the theta
368
+ device (str): the device to store the precomputed cis
369
+ """
370
+
371
+ def __init__(self, dim: int, max_height: int, max_width: int, theta_base=10000):
372
+ super().__init__()
373
+ self.dim = dim
374
+ assert self.dim % 4 == 0, "dim must be divisible by 4"
375
+ self.max_height = max_height
376
+ self.max_width = max_width
377
+ self.theta_base = theta_base
378
+
379
+ def extra_repr(self):
380
+ return f"dim={self.dim}, max_height={self.max_height}, max_width={self.max_width}, theta_base={self.theta_base}"
381
+
382
+ def _precompute_freqs_cis(self, device: torch.device) -> torch.Tensor:
383
+ """Calculate the cis(freqs) for each position in the 2D grid.
384
+
385
+ Return: complex tensor of shape (max_height, max_width, dim//2) and value:
386
+ height axis: ret[h, w, 2*i] = cis(h * theta_base**(-4*i/dim))
387
+ weight axis: ret[h, w, 2*i+1] = cis(w * theta_base**(-4*i/dim)) with (i in [0, dim//4))
388
+ note: `cis` is a mathematical notation defined by cis x = cos x + i sin x,
389
+ """
390
+ N = self.max_height * self.max_width
391
+ flat_pos = torch.arange(0, N).float().to(device)
392
+ x_pos = flat_pos % self.max_width
393
+ y_pos = flat_pos // self.max_width
394
+ dim_range = (
395
+ torch.arange(0, self.dim, 4)[: (self.dim // 4)].float().to(device)
396
+ ) # C/4
397
+ freqs = 1.0 / (self.theta_base ** (dim_range / self.dim))
398
+ x_freqs = torch.outer(x_pos, freqs).float() # N, C/4
399
+ y_freqs = torch.outer(y_pos, freqs).float() # N, C/4
400
+ x_cis = torch.polar(torch.ones_like(x_freqs), x_freqs) # N, C/4
401
+ y_cis = torch.polar(torch.ones_like(y_freqs), y_freqs) # N, C/4
402
+ # N, C/4, 2
403
+ freqs_cis = torch.cat(
404
+ [x_cis.unsqueeze(dim=-1), y_cis.unsqueeze(dim=-1)], dim=-1
405
+ )
406
+ # max_height, max_width, C/2
407
+ freqs_cis = freqs_cis.reshape(self.max_height, self.max_width, -1)
408
+ return freqs_cis
409
+
410
+ def get_freqs_cis(
411
+ self, grid_thws: torch.Tensor, device: torch.device
412
+ ) -> torch.Tensor:
413
+ """
414
+ Args:
415
+ grid_thws (torch.Tensor): grid time, height and width
416
+
417
+ Returns:
418
+ freqs_cis: tensor of shape (sum(t * height * width), dim//2)
419
+ """
420
+ if not hasattr(self, "freqs_cis"):
421
+ self.register_buffer(
422
+ "freqs_cis", self._precompute_freqs_cis(device), persistent=False
423
+ )
424
+
425
+ shapes = grid_thws.tolist()
426
+ assert all(
427
+ 1 <= h <= self.max_height and 1 <= w <= self.max_width for t, h, w in shapes
428
+ ), (
429
+ shapes,
430
+ self.max_height,
431
+ self.max_width,
432
+ )
433
+ freqs_cis = torch.cat(
434
+ [
435
+ self.freqs_cis[:h, :w].reshape(-1, self.dim // 2).repeat(t, 1)
436
+ for t, h, w in shapes
437
+ ],
438
+ dim=0,
439
+ )
440
+ return freqs_cis
441
+
442
+
443
+ class MLP2(nn.Module):
444
+ """
445
+ Args:
446
+ dims: [in_dim, hidden_dim, out_dim]
447
+ bias: whether to use bias in linear layer.
448
+ """
449
+
450
+ def __init__(self, dims: list[int], activation, bias=True):
451
+ super().__init__()
452
+ assert len(dims) == 3
453
+ self.fc0 = nn.Linear(dims[0], dims[1], bias=bias)
454
+ self.fc1 = nn.Linear(dims[1], dims[2], bias=bias)
455
+ self.activation = activation
456
+ for m in [self.fc0, self.fc1]:
457
+ nn.init.trunc_normal_(m.weight, std=math.sqrt(2 / m.in_features))
458
+ if m.bias is not None:
459
+ nn.init.zeros_(m.bias)
460
+
461
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
462
+ x = self.fc0(x)
463
+ x = self.activation(x)
464
+ return self.fc1(x)
465
+
466
+
467
+ class MoonViTEncoderLayer(nn.Module):
468
+ def __init__(
469
+ self,
470
+ num_heads: int,
471
+ hidden_dim: int,
472
+ mlp_dim: int,
473
+ qkv_hidden_size: int | None = None,
474
+ norm_type: str = "layernorm",
475
+ mlp_type: str = "mlp2",
476
+ *,
477
+ attn_implementation: str = "flash_attention_2",
478
+ activation=F.gelu,
479
+ attn_bias: bool = False,
480
+ linear_bias: bool = True,
481
+ use_deterministic_attn: bool = False,
482
+ ):
483
+ super().__init__()
484
+ self.num_heads = num_heads
485
+ self.hidden_dim = hidden_dim
486
+ self.qkv_hidden_size = (
487
+ hidden_dim if qkv_hidden_size is None else qkv_hidden_size
488
+ )
489
+ self.hidden_size_per_attention_head = self.qkv_hidden_size // self.num_heads
490
+ self.attn_implementation = attn_implementation
491
+ self.use_deterministic_attn = use_deterministic_attn
492
+
493
+ if norm_type == "layernorm":
494
+ self.norm0 = nn.LayerNorm(hidden_dim)
495
+ self.norm1 = nn.LayerNorm(hidden_dim)
496
+ elif norm_type == "rmsnorm":
497
+ self.norm0 = nn.RMSNorm(hidden_dim)
498
+ self.norm1 = nn.RMSNorm(hidden_dim)
499
+ else:
500
+ raise NotImplementedError(f"Not support norm_type: {norm_type}")
501
+
502
+ if mlp_type == "mlp2":
503
+ self.mlp = MLP2(
504
+ [hidden_dim, mlp_dim, hidden_dim], activation, bias=linear_bias
505
+ )
506
+ else:
507
+ raise NotImplementedError(f"Not support mlp_type: {mlp_type}")
508
+
509
+ self.wqkv = nn.Linear(hidden_dim, self.qkv_hidden_size * 3, bias=attn_bias)
510
+ self.wo = nn.Linear(self.qkv_hidden_size, hidden_dim, bias=attn_bias)
511
+
512
+ def attention_qkvpacked(
513
+ self,
514
+ x: torch.Tensor,
515
+ cu_seqlens: torch.Tensor,
516
+ max_seqlen: torch.Tensor,
517
+ rope_freqs_cis: torch.Tensor | None = None,
518
+ ):
519
+ """
520
+ Args:
521
+ x (torch.Tensor): (batch_size, seqlen, hidden_dim)
522
+ cu_seqlens (torch.Tensor):
523
+ """
524
+ xqkv = self.wqkv(x)
525
+
526
+ qkv_shape = xqkv.size()[:-1] + (
527
+ 3,
528
+ self.num_heads,
529
+ self.hidden_size_per_attention_head,
530
+ )
531
+ # xqkv: (batch_size, seqlen, 3, nheads, headdim)
532
+ xqkv = xqkv.view(*qkv_shape)
533
+ xq, xk, xv = torch.unbind(xqkv, dim=-3)
534
+
535
+ xq, xk = apply_rope(xq, xk, rope_freqs_cis)
536
+
537
+ attn_func = VL_VISION_ATTENTION_FUNCTIONS[self.attn_implementation]
538
+ attn_out = attn_func(
539
+ xq,
540
+ xk,
541
+ xv,
542
+ q_cu_seqlens=cu_seqlens,
543
+ k_cu_seqlens=cu_seqlens,
544
+ max_seqlen_k=max_seqlen,
545
+ max_seqlen_q=max_seqlen,
546
+ deterministic=self.use_deterministic_attn,
547
+ )
548
+
549
+ attn_out = self.wo(attn_out)
550
+ return attn_out
551
+
552
+ def forward(
553
+ self,
554
+ hidden_states: torch.Tensor,
555
+ cu_seqlens: torch.Tensor,
556
+ max_seqlen: int,
557
+ rope_freqs_cis: torch.Tensor | None = None,
558
+ ):
559
+ residual = hidden_states
560
+ hidden_states = self.norm0(hidden_states)
561
+
562
+ hidden_states = self.attention_qkvpacked(
563
+ hidden_states, cu_seqlens, max_seqlen, rope_freqs_cis
564
+ )
565
+ hidden_states = residual + hidden_states
566
+
567
+ residual = hidden_states
568
+ hidden_states = self.norm1(hidden_states)
569
+ hidden_states = self.mlp(hidden_states)
570
+ hidden_states = residual + hidden_states
571
+
572
+ return hidden_states
573
+
574
+
575
+ class MoonViT3dEncoder(nn.Module):
576
+ def __init__(
577
+ self,
578
+ hidden_dim: int,
579
+ num_layers: int,
580
+ block_cfg: dict,
581
+ use_deterministic_attn: bool = False,
582
+ ) -> None:
583
+ super().__init__()
584
+ self.use_deterministic_attn = use_deterministic_attn
585
+
586
+ qkv_hidden_size = (
587
+ block_cfg["hidden_dim"]
588
+ if block_cfg.get("qkv_hidden_size") is None
589
+ else block_cfg["qkv_hidden_size"]
590
+ )
591
+ self.rope_2d = Rope2DPosEmbRepeated(
592
+ qkv_hidden_size // block_cfg["num_heads"], 512, 512
593
+ )
594
+ self.blocks = nn.ModuleList(
595
+ [
596
+ MoonViTEncoderLayer(
597
+ **block_cfg, use_deterministic_attn=self.use_deterministic_attn
598
+ )
599
+ for _ in range(num_layers)
600
+ ]
601
+ )
602
+ norm_type = block_cfg.get("norm_type", "layernorm")
603
+ if norm_type == "layernorm":
604
+ self.final_layernorm = nn.LayerNorm(hidden_dim)
605
+ elif norm_type == "rmsnorm":
606
+ self.final_layernorm = nn.RMSNorm(hidden_dim)
607
+ else:
608
+ raise NotImplementedError(f"Not support norm_type: {norm_type}")
609
+
610
+ def forward(
611
+ self,
612
+ hidden_states: torch.Tensor,
613
+ grid_thws: torch.Tensor,
614
+ ) -> torch.Tensor:
615
+ rope_freqs_cis = self.rope_2d.get_freqs_cis(
616
+ grid_thws=grid_thws, device=hidden_states.device
617
+ )
618
+
619
+ lengths = torch.cat(
620
+ (
621
+ torch.zeros(1, dtype=grid_thws.dtype, device=grid_thws.device),
622
+ grid_thws[:, 0] * grid_thws[:, 1] * grid_thws[:, 2],
623
+ )
624
+ )
625
+
626
+ max_seqlen = lengths.max()
627
+ cu_seqlens = lengths.to(hidden_states.device).cumsum(dim=0, dtype=torch.int32)
628
+ for block in self.blocks:
629
+ hidden_states = block(
630
+ hidden_states, cu_seqlens, max_seqlen, rope_freqs_cis=rope_freqs_cis
631
+ )
632
+
633
+ hidden_states = self.final_layernorm(hidden_states)
634
+ return hidden_states
635
+
636
+
637
+ def tpool_patch_merger(
638
+ x: torch.Tensor,
639
+ grid_thws: torch.Tensor,
640
+ merge_kernel_size: tuple[int, int] = (2, 2),
641
+ ) -> list[torch.Tensor]:
642
+ d_model = x.size(-1)
643
+
644
+ outputs = []
645
+ pre_sum = 0
646
+ for t, h, w in grid_thws.tolist():
647
+ # Get the current sequence
648
+ seq = x[pre_sum : pre_sum + t * h * w]
649
+ # Reshape along self.merge_kernel_size and concat to the last dimension
650
+ kernel_height, kernel_width = merge_kernel_size
651
+ new_height, new_width = h // kernel_height, w // kernel_width
652
+ reshaped_seq = seq.view(
653
+ t, new_height, kernel_height, new_width, kernel_width, d_model
654
+ )
655
+ reshaped_seq = (
656
+ reshaped_seq.permute(0, 1, 3, 2, 4, 5).contiguous().mean(dim=0)
657
+ ) # temporal pooling
658
+ padded_seq = reshaped_seq.view(
659
+ new_height * new_width, kernel_height * kernel_width, -1
660
+ )
661
+ outputs.append(padded_seq)
662
+ pre_sum += t * h * w
663
+
664
+ return outputs
665
+
666
+
667
+ class MoonViT3dPretrainedModel(PreTrainedModel):
668
+ config_class = None
669
+ model_type = "moonvit3d"
670
+ _no_split_modules = ["MoonViTEncoderLayer"]
671
+ _supports_flash_attn_2 = True
672
+ _supports_sdpa = True
673
+
674
+ def __init__(self, config, *inputs, **kwargs):
675
+ super().__init__(config, *inputs, **kwargs)
676
+ config = deepcopy(config)
677
+ self.merge_kernel_size = config.merge_kernel_size
678
+ self.patch_size = config.patch_size
679
+ self.merge_type = config.merge_type
680
+
681
+ self.patch_embed = MoonVision3dPatchEmbed(
682
+ out_dim=config.hidden_size,
683
+ patch_size=config.patch_size,
684
+ pos_emb_height=config.init_pos_emb_height,
685
+ pos_emb_width=config.init_pos_emb_width,
686
+ pos_emb_time=config.init_pos_emb_time,
687
+ pos_emb_type=config.pos_emb_type,
688
+ patch_embed_proj_bias=getattr(config, "patch_embed_proj_bias", True),
689
+ pos_emb_interpolation_mode=getattr(
690
+ config, "pos_emb_interpolation_mode", "bicubic"
691
+ ),
692
+ )
693
+
694
+ self.encoder = MoonViT3dEncoder(
695
+ hidden_dim=config.hidden_size,
696
+ num_layers=config.num_hidden_layers,
697
+ block_cfg={
698
+ "num_heads": config.num_attention_heads,
699
+ "hidden_dim": config.hidden_size,
700
+ "qkv_hidden_size": getattr(config, "qkv_hidden_size", None),
701
+ "mlp_dim": config.intermediate_size,
702
+ "norm_type": getattr(config, "norm_type", "layernorm"),
703
+ "mlp_type": getattr(config, "mlp_type", "mlp2"),
704
+ "activation": PytorchGELUTanh(),
705
+ "attn_bias": getattr(config, "attn_bias", True),
706
+ "linear_bias": getattr(config, "linear_bias", True),
707
+ "attn_implementation": config._attn_implementation,
708
+ },
709
+ use_deterministic_attn=getattr(self, "use_deterministic_attn", False),
710
+ )
711
+
712
+ def forward(
713
+ self, pixel_values: torch.Tensor, grid_thws: torch.Tensor
714
+ ) -> torch.Tensor:
715
+ """
716
+ Args:
717
+ pixel_values (torch.Tensor): The input pixel values.
718
+ grid_thws (torch.Tensor): Temporal, height and width.
719
+
720
+ Returns:
721
+ torch.Tensor: The output tokens.
722
+ """
723
+ # grid_thws = grid_thws.to('cpu')
724
+ assert grid_thws.ndim == 2, f"grid_thws should be 2D, got {grid_thws.ndim}"
725
+ assert grid_thws.size(1) == 3, f"No support for thw: {grid_thws}"
726
+ hidden_states = self.patch_embed(pixel_values, grid_thws)
727
+ hidden_states = self.encoder(hidden_states, grid_thws)
728
+ if (
729
+ self.merge_type == "sd2_tpool"
730
+ ): # spatial downsampling 2x with temporal pooling all
731
+ hidden_states = tpool_patch_merger(
732
+ hidden_states, grid_thws, merge_kernel_size=self.merge_kernel_size
733
+ )
734
+ else:
735
+ raise NotImplementedError(f"Not support {self.merge_type}")
736
+
737
+ return hidden_states
738
+
739
+
740
+ # ============================================================================
741
+ # MM Projector Helper Classes (from mm_projector/modeling_mm_projectors.py)
742
+ # ============================================================================
743
+
744
+
745
+ class IdentityMap(nn.Module):
746
+ def __init__(self):
747
+ super().__init__()
748
+
749
+ def forward(self, x, *args, **kwargs):
750
+ return x
751
+
752
+
753
+ class MLP(nn.Module):
754
+ def __init__(self, config):
755
+ super().__init__()
756
+ # TODO, use faster LayerNorm
757
+ self.pre_norm = nn.LayerNorm(config.mm_hidden_size)
758
+ self.proj = nn.Sequential(
759
+ nn.Linear(config.mm_hidden_size, config.hidden_size),
760
+ nn.GELU(),
761
+ nn.Linear(config.hidden_size, config.hidden_size),
762
+ )
763
+
764
+ def forward(self, x, *args, **kwargs):
765
+ assert isinstance(x, list | tuple), f"x is not a list or tuple: {type(x)}"
766
+ lengths = [item.shape[0] for item in x]
767
+ x = torch.cat(x, dim=0)
768
+ x = self.pre_norm(x)
769
+ x = self.proj(x)
770
+ x = torch.split(x, lengths, dim=0)
771
+
772
+ return x
773
+
774
+
775
+ class PatchMergerMLP(nn.Module):
776
+ def __init__(self, config):
777
+ super().__init__()
778
+ eps = config.projector_ln_eps
779
+ self.hidden_size = config.mm_hidden_size * (
780
+ config.merge_kernel_size[0] * config.merge_kernel_size[1]
781
+ )
782
+ self.pre_norm = nn.LayerNorm(config.mm_hidden_size, eps=eps)
783
+ self.proj = nn.Sequential(
784
+ nn.Linear(self.hidden_size, self.hidden_size),
785
+ nn.GELU(),
786
+ nn.Linear(self.hidden_size, config.hidden_size),
787
+ )
788
+
789
+ def forward(self, x, *args, **kwargs):
790
+ if isinstance(x, list) or isinstance(x, tuple):
791
+ x = [self.proj(self.pre_norm(item).view(item.shape[0], -1)) for item in x]
792
+ else:
793
+ # B, N, N_k, C = x.shape
794
+ B = x.shape[0]
795
+ x = self.proj(self.pre_norm(x).view(B, -1, self.hidden_size))
796
+ return x
797
+
798
+
799
+ class PatchMergerMLPV2(nn.Module):
800
+ def __init__(self, config):
801
+ super().__init__()
802
+ eps = config.projector_ln_eps
803
+ self.hidden_size = config.mm_hidden_size * (
804
+ config.merge_kernel_size[0] * config.merge_kernel_size[1]
805
+ )
806
+ self.proj = nn.Sequential(
807
+ nn.Linear(self.hidden_size, self.hidden_size, bias=False),
808
+ nn.GELU(),
809
+ nn.Linear(self.hidden_size, config.hidden_size, bias=False),
810
+ )
811
+ self.post_norm = nn.RMSNorm(config.hidden_size, eps=eps)
812
+ for m in self.proj.modules():
813
+ if isinstance(m, nn.Linear):
814
+ nn.init.trunc_normal_(m.weight, std=math.sqrt(2 / m.in_features))
815
+ if m.bias is not None:
816
+ nn.init.zeros_(m.bias)
817
+
818
+ def forward(self, x, *args, **kwargs):
819
+ if isinstance(x, list) or isinstance(x, tuple):
820
+ lengths = [item.shape[0] for item in x]
821
+ x = torch.concat([item.view(item.shape[0], -1) for item in x], dim=0)
822
+ x = self.post_norm(self.proj(x))
823
+ x = torch.split(x, lengths, dim=0)
824
+ else:
825
+ # B, N, N_k, C = x.shape
826
+ B = x.shape[0]
827
+ x = self.proj(x.view(B, -1, self.hidden_size))
828
+ x = self.post_norm(x)
829
+ return x
830
+
831
+
832
+ class KimiK3PreTrainedModel(PreTrainedModel):
833
+ config_class = KimiK3Config
834
+ base_model_prefix = "model"
835
+ _no_split_modules = [
836
+ "MoonViT3dPretrainedModel",
837
+ "MoonViTEncoderLayer",
838
+ "KimiDecoderLayer",
839
+ "PatchMergerMLP",
840
+ "PatchMergerMLPV2",
841
+ ]
842
+ _skip_keys_device_placement = "past_key_values"
843
+ _supports_flash_attn_2 = True
844
+ _supports_sdpa = False
845
+
846
+ def _init_weights(self, module):
847
+ # HOTFIX: disk offloading attempts to initialize the meta tensors
848
+ # but this is bad programming: we shouldn't be initializing these
849
+ # params in the first place
850
+ # the init attempt attempts to get `module.weight`, which DNE for qmodels
851
+
852
+ return
853
+
854
+ # important: this ported version of Llava isn't meant for training from scratch - only
855
+ # inference and fine-tuning - so the proper init weights code has been removed - the original codebase
856
+ # https://github.com/haotian-liu/LLaVA/tree/main/llava should serve for that purpose
857
+ std = (
858
+ self.config.initializer_range
859
+ if hasattr(self.config, "initializer_range")
860
+ else self.config.text_config.initializer_range
861
+ )
862
+
863
+ if hasattr(module, "class_embedding"):
864
+ module.class_embedding.data.normal_(mean=0.0, std=std)
865
+
866
+ if isinstance(module, (nn.Linear, nn.Conv2d)):
867
+ module.weight.data.normal_(mean=0.0, std=std)
868
+ if module.bias is not None:
869
+ module.bias.data.zero_()
870
+ elif isinstance(module, nn.Embedding):
871
+ module.weight.data.normal_(mean=0.0, std=std)
872
+ if module.padding_idx is not None:
873
+ module.weight.data[module.padding_idx].zero_()
874
+
875
+
876
+ class VisionTowerConfig(PretrainedConfig):
877
+ model_type = "moonvit3d"
878
+
879
+ def __init__(self, config: KimiK3Config, **kwargs):
880
+ super().__init__(**kwargs)
881
+ self.patch_size = config.patch_size
882
+ self.init_pos_emb_height = config.init_pos_emb_height
883
+ self.init_pos_emb_width = config.init_pos_emb_width
884
+ self.init_pos_emb_time = config.init_pos_emb_time
885
+ self.pos_emb_type = config.pos_emb_type
886
+ self.num_attention_heads = config.vt_num_attention_heads
887
+ self.num_hidden_layers = config.vt_num_hidden_layers
888
+ self.hidden_size = config.vt_hidden_size
889
+ self.intermediate_size = config.vt_intermediate_size
890
+ self.merge_kernel_size = config.merge_kernel_size
891
+ self.merge_type = config.merge_type
892
+ self._attn_implementation = config._attn_implementation
893
+ self.qkv_hidden_size = getattr(config, "qkv_hidden_size", None)
894
+ self.norm_type = getattr(config, "norm_type", "layernorm")
895
+ self.attn_bias = getattr(config, "attn_bias", True)
896
+ self.patch_embed_proj_bias = getattr(config, "patch_embed_proj_bias", True)
897
+ self.mlp_type = getattr(config, "mlp_type", "mlp2")
898
+ self.linear_bias = getattr(config, "linear_bias", True)
899
+ self.pos_emb_interpolation_mode = getattr(
900
+ config, "pos_emb_interpolation_mode", "bilinear"
901
+ )
902
+
903
+
904
+ class ProjectorConfig:
905
+ def __init__(self, config: KimiK3Config):
906
+ self.mm_projector_type = config.mm_projector_type
907
+ self.mm_hidden_size = config.mm_hidden_size
908
+ self.hidden_size = config.text_hidden_size
909
+ self.merge_kernel_size = config.merge_kernel_size
910
+ self.projector_hidden_act = config.projector_hidden_act
911
+ self.projector_ln_eps = config.projector_ln_eps
912
+
913
+
914
+ # ref https://github.com/huggingface/transformers/blob/78b2929c0554b79e0489b451ce4ece14d265ead2/src/transformers/models/llava/modeling_llava.py#L240
915
+ class KimiK3ForConditionalGeneration(KimiK3PreTrainedModel):
916
+ @classmethod
917
+ def _supports_default_dynamic_cache(cls) -> bool:
918
+ return False
919
+
920
+ def __init__(self, config: KimiK3Config):
921
+ super().__init__(config)
922
+
923
+ vt_config = VisionTowerConfig(config.vision_config)
924
+ self.vision_tower = MoonViT3dPretrainedModel(vt_config)
925
+
926
+ proj_config = ProjectorConfig(config.vision_config)
927
+ if proj_config.mm_projector_type == "identity":
928
+ self.mm_projector = IdentityMap()
929
+ elif proj_config.mm_projector_type == "mlp":
930
+ self.mm_projector = MLP(proj_config)
931
+ elif proj_config.mm_projector_type == "patchmerger":
932
+ self.mm_projector = PatchMergerMLP(proj_config)
933
+ elif proj_config.mm_projector_type == "patchmergerv2":
934
+ self.mm_projector = PatchMergerMLPV2(proj_config)
935
+ else:
936
+ raise ValueError(
937
+ f"Unsupported mm_projector_type: {proj_config.mm_projector_type}"
938
+ )
939
+
940
+ self.language_model = KimiLinearForCausalLM(config.text_config)
941
+ self.post_init()
942
+
943
+ if hasattr(self.language_model, "dtype"):
944
+ target_dtype = self.language_model.dtype
945
+ self.vision_tower = self.vision_tower.to(dtype=target_dtype)
946
+ self.mm_projector = self.mm_projector.to(dtype=target_dtype)
947
+
948
+ def get_input_embeddings(self):
949
+ return self.language_model.get_input_embeddings()
950
+
951
+ def set_input_embeddings(self, value):
952
+ self.language_model.set_input_embeddings(value)
953
+
954
+ def get_output_embeddings(self):
955
+ return self.language_model.get_output_embeddings()
956
+
957
+ def set_output_embeddings(self, new_embeddings):
958
+ self.language_model.set_output_embeddings(new_embeddings)
959
+
960
+ def set_decoder(self, decoder):
961
+ self.language_model.set_decoder(decoder)
962
+
963
+ def get_decoder(self):
964
+ return self.language_model.get_decoder()
965
+
966
+ def tie_weights(self, **kwargs):
967
+ return self.language_model.tie_weights(**kwargs)
968
+
969
+ def resize_token_embeddings(
970
+ self, new_num_tokens: int | None = None, pad_to_multiple_of=None
971
+ ) -> nn.Embedding:
972
+ model_embeds = self.language_model.resize_token_embeddings(
973
+ new_num_tokens, pad_to_multiple_of
974
+ )
975
+ # update vocab size
976
+ self.config.text_config.vocab_size = model_embeds.num_embeddings
977
+ self.vocab_size = model_embeds.num_embeddings
978
+ return model_embeds
979
+
980
+ def _merge_input_ids_with_image_features(
981
+ self,
982
+ image_features: list[torch.Tensor],
983
+ inputs_embeds: torch.Tensor,
984
+ input_ids: torch.Tensor,
985
+ attention_mask: torch.Tensor,
986
+ labels: torch.Tensor | None = None,
987
+ ):
988
+ """
989
+ Args:
990
+ image_features (:obj:`torch.Tensor` of shape :obj:`(num_image_tokens, embed_dim)`):
991
+ The image features to merge with the input embeddings.
992
+ inputs_embeds (:obj:`torch.Tensor` of shape :obj:`(batch_size, sequence_length, embed_dim)`):
993
+ The input embeddings.
994
+ input_ids (:obj:`torch.Tensor` of shape :obj:`(batch_size, sequence_length)`):
995
+ The input ids.
996
+ attention_mask (:obj:`torch.Tensor` of shape :obj:`(batch_size, sequence_length)`):
997
+ The attention mask.
998
+ labels (:obj:`torch.Tensor` of shape :obj:`(batch_size, sequence_length)`, *optional*):
999
+ The labels.
1000
+ """
1001
+ _, embed_dim = image_features[0].shape
1002
+ feature_lengths = [x.shape[0] for x in image_features]
1003
+ image_features = torch.cat(image_features, dim=0)
1004
+
1005
+ image_token_index: int = self.config.media_placeholder_token_id
1006
+ pad_token_id: int = self.config.pad_token_id
1007
+ ignore_index: int = self.config.ignore_index
1008
+
1009
+ batch_size, sequence_length = input_ids.shape
1010
+ left_padding = not torch.sum(input_ids[:, -1] == torch.tensor(pad_token_id))
1011
+
1012
+ # 1. Create a mask to know where special image tokens are
1013
+ _token_occupation_table = torch.ones_like(input_ids.flatten())
1014
+ _token_occupation_table[input_ids.flatten() == image_token_index] = (
1015
+ torch.tensor(feature_lengths, dtype=torch.long, device=input_ids.device)
1016
+ )
1017
+ _token_occupation_table = _token_occupation_table.reshape(input_ids.shape)
1018
+
1019
+ max_embed_dim = _token_occupation_table.sum(-1).max().item()
1020
+ assert (
1021
+ max_embed_dim >= sequence_length
1022
+ ), f"The maximum embedding dimension ({max_embed_dim}) is less than the sequence length ({sequence_length})"
1023
+ batch_indices, non_image_indices = torch.where(input_ids != image_token_index)
1024
+
1025
+ # 2. Compute the positions where text should be written
1026
+ # Calculate new positions for text tokens in merged image-text sequence.
1027
+ new_token_positions = torch.cumsum(_token_occupation_table, -1) - 1
1028
+ nb_image_pad = max_embed_dim - 1 - new_token_positions[:, -1]
1029
+ if left_padding:
1030
+ new_token_positions += nb_image_pad[:, None] # offset for left padding
1031
+ text_to_overwrite = new_token_positions[batch_indices, non_image_indices]
1032
+
1033
+ # 3. Create the full embedding, already padded to the maximum position
1034
+ final_embedding = torch.zeros(
1035
+ batch_size,
1036
+ max_embed_dim,
1037
+ embed_dim,
1038
+ dtype=inputs_embeds.dtype,
1039
+ device=inputs_embeds.device,
1040
+ )
1041
+ final_attention_mask = torch.zeros(
1042
+ batch_size,
1043
+ max_embed_dim,
1044
+ dtype=attention_mask.dtype,
1045
+ device=inputs_embeds.device,
1046
+ )
1047
+ if labels is not None:
1048
+ final_labels = torch.full(
1049
+ (batch_size, max_embed_dim),
1050
+ ignore_index,
1051
+ dtype=input_ids.dtype,
1052
+ device=input_ids.device,
1053
+ )
1054
+ # In case the Vision model or the Language model has been offloaded to CPU, we need to manually
1055
+ # set the corresponding tensors into their correct target device.
1056
+ target_device = inputs_embeds.device
1057
+ batch_indices, non_image_indices, text_to_overwrite = (
1058
+ batch_indices.to(target_device),
1059
+ non_image_indices.to(target_device),
1060
+ text_to_overwrite.to(target_device),
1061
+ )
1062
+ attention_mask = attention_mask.to(target_device)
1063
+
1064
+ # 4. Fill the embeddings based on the mask.
1065
+ final_embedding[batch_indices, text_to_overwrite] = inputs_embeds[
1066
+ batch_indices, non_image_indices
1067
+ ]
1068
+ final_attention_mask[batch_indices, text_to_overwrite] = attention_mask[
1069
+ batch_indices, non_image_indices
1070
+ ]
1071
+ if labels is not None:
1072
+ final_labels[batch_indices, text_to_overwrite] = labels[
1073
+ batch_indices, non_image_indices
1074
+ ]
1075
+
1076
+ # 5. Fill the embeddings corresponding to the images. Anything that is not `text_positions` needs filling (#29835)
1077
+ image_to_overwrite = torch.full(
1078
+ (batch_size, max_embed_dim),
1079
+ True,
1080
+ dtype=torch.bool,
1081
+ device=inputs_embeds.device,
1082
+ )
1083
+ image_to_overwrite[batch_indices, text_to_overwrite] = False
1084
+ image_to_overwrite &= image_to_overwrite.cumsum(-1) - 1 >= nb_image_pad[
1085
+ :, None
1086
+ ].to(target_device)
1087
+
1088
+ if image_to_overwrite.sum() != image_features.shape[:-1].numel():
1089
+ raise ValueError(
1090
+ f"The input provided to the model are wrong. The number of image tokens is {image_to_overwrite.sum()} while"
1091
+ f" the number of image features given to the model is {image_features.shape[:-1].numel()}. "
1092
+ "This prevents correct indexing and breaks batch generation."
1093
+ )
1094
+
1095
+ final_embedding[image_to_overwrite] = (
1096
+ image_features.contiguous().reshape(-1, embed_dim).to(target_device)
1097
+ )
1098
+ final_attention_mask |= image_to_overwrite
1099
+ position_ids = (final_attention_mask.cumsum(-1) - 1).masked_fill_(
1100
+ (final_attention_mask == 0), 1
1101
+ )
1102
+
1103
+ # 6. Mask out the embedding at padding positions, as we later use the past_key_value value to determine the non-attended tokens.
1104
+ batch_indices, pad_indices = torch.where(input_ids == pad_token_id)
1105
+ indices_to_mask = new_token_positions[batch_indices, pad_indices]
1106
+
1107
+ final_embedding[batch_indices, indices_to_mask] = 0
1108
+
1109
+ if labels is None:
1110
+ final_labels = None
1111
+
1112
+ return final_embedding, final_attention_mask, final_labels, position_ids
1113
+
1114
+ def _extract_image_features(
1115
+ self, pixel_values: torch.Tensor, grid_thws: torch.Tensor
1116
+ ) -> list[torch.Tensor]:
1117
+ """
1118
+ Args:
1119
+ pixel_values (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, num_channels, height, width)`):
1120
+ The pixel values of the images processed by image processor.
1121
+ grid_thws (:obj:`torch.Tensor` of shape :obj:`(batch_size, 3)`):
1122
+ The grid, height, width of the images.
1123
+
1124
+ Returns:
1125
+ selected_image_feature (:obj:`torch.FloatTensor` of shape :obj:`(num_image_tokens, embed_dim)`):
1126
+ The selected image features to use as input to the projector head.
1127
+
1128
+ """
1129
+
1130
+ target_dtype = self.vision_tower.patch_embed.proj.weight.dtype
1131
+ pixel_values = pixel_values.to(target_dtype)
1132
+
1133
+ image_features = self.vision_tower(pixel_values, grid_thws)
1134
+ return image_features
1135
+
1136
+ def forward(
1137
+ self,
1138
+ input_ids: torch.LongTensor | None = None,
1139
+ pixel_values: torch.FloatTensor | list[torch.FloatTensor] | None = None,
1140
+ grid_thws: torch.Tensor | None = None,
1141
+ attention_mask: torch.Tensor | None = None,
1142
+ position_ids: torch.LongTensor | None = None,
1143
+ past_key_values: list[torch.FloatTensor] | None = None,
1144
+ inputs_embeds: torch.FloatTensor | None = None,
1145
+ labels: torch.LongTensor | None = None,
1146
+ use_cache: bool | None = None,
1147
+ output_attentions: bool | None = None,
1148
+ output_hidden_states: bool | None = None,
1149
+ return_dict: bool | None = None,
1150
+ ) -> tuple | LlavaCausalLMOutputWithPast:
1151
+ r"""
1152
+ Args:
1153
+ labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
1154
+ Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
1155
+ config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
1156
+ (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
1157
+
1158
+ ```"""
1159
+ assert self.vision_tower is not None, "vision_tower is not loaded"
1160
+ output_attentions = (
1161
+ output_attentions
1162
+ if output_attentions is not None
1163
+ else self.config.output_attentions
1164
+ )
1165
+ output_hidden_states = (
1166
+ output_hidden_states
1167
+ if output_hidden_states is not None
1168
+ else self.config.output_hidden_states
1169
+ )
1170
+ return_dict = (
1171
+ return_dict if return_dict is not None else self.config.use_return_dict
1172
+ )
1173
+
1174
+ if inputs_embeds is None:
1175
+ # 1. Extra the input embeddings
1176
+ inputs_embeds = self.get_input_embeddings()(input_ids)
1177
+
1178
+ # 2. Merge text and images
1179
+ if (
1180
+ pixel_values is not None
1181
+ and len(pixel_values) > 0
1182
+ and input_ids.shape[1] != 1
1183
+ ):
1184
+ image_features = self._extract_image_features(pixel_values, grid_thws)
1185
+ if self.mm_projector:
1186
+ image_features = self.mm_projector(image_features)
1187
+
1188
+ inputs_embeds = inputs_embeds.to(
1189
+ image_features[0].dtype
1190
+ ) # num_tokens, embed_dim
1191
+ inputs_embeds, attention_mask, labels, position_ids = (
1192
+ self._merge_input_ids_with_image_features(
1193
+ image_features,
1194
+ inputs_embeds,
1195
+ input_ids,
1196
+ attention_mask,
1197
+ labels,
1198
+ )
1199
+ )
1200
+
1201
+ # In case input_ids.shape[1] == 1 & pixel_values==None & past_key_values != None, we are in the case of
1202
+ # generation with cache
1203
+ elif (
1204
+ past_key_values is not None
1205
+ and pixel_values is not None
1206
+ and input_ids.shape[1] == 1
1207
+ ):
1208
+ # Retrieve the first layer to inspect the logits and mask out the hidden states
1209
+ # that are set to 0
1210
+ first_layer_past_key_value = past_key_values[0][0][:, :, :, 0]
1211
+
1212
+ # Sum all dimensions of head_dim (-2) to avoid random errors such as: https://github.com/huggingface/transformers/pull/28032#issuecomment-1863691941
1213
+ batch_index, non_attended_tokens = torch.where(
1214
+ first_layer_past_key_value.float().sum(-2) == 0
1215
+ )
1216
+
1217
+ # Get the target length
1218
+ target_length = input_ids.shape[1]
1219
+ past_length = first_layer_past_key_value.shape[-1]
1220
+
1221
+ extended_attention_mask = torch.ones(
1222
+ (attention_mask.shape[0], past_length),
1223
+ dtype=attention_mask.dtype,
1224
+ device=attention_mask.device,
1225
+ )
1226
+
1227
+ # Filter out only the tokens that can be un-attended, this can happen
1228
+ # if one uses Llava + Fused modules where the cache on the
1229
+ # first iteration is already big enough, or if one passes custom cache
1230
+ valid_indices = non_attended_tokens < extended_attention_mask.size(-1)
1231
+ new_batch_index = batch_index[valid_indices]
1232
+ new_non_attended_tokens = non_attended_tokens[valid_indices]
1233
+
1234
+ # Zero-out the places where we don't need to attend
1235
+ extended_attention_mask[new_batch_index, new_non_attended_tokens] = 0
1236
+
1237
+ attention_mask = torch.cat(
1238
+ (extended_attention_mask, attention_mask[:, -target_length:]), dim=1
1239
+ )
1240
+ position_ids = torch.sum(attention_mask, dim=1).unsqueeze(-1) - 1
1241
+
1242
+ outputs = self.language_model(
1243
+ attention_mask=attention_mask,
1244
+ position_ids=position_ids,
1245
+ past_key_values=past_key_values,
1246
+ inputs_embeds=inputs_embeds,
1247
+ use_cache=use_cache,
1248
+ output_attentions=output_attentions,
1249
+ output_hidden_states=output_hidden_states,
1250
+ return_dict=return_dict,
1251
+ )
1252
+
1253
+ logits = outputs[0]
1254
+
1255
+ loss = None
1256
+ if labels is not None:
1257
+ # Shift so that tokens < n predict n
1258
+ if attention_mask is not None:
1259
+ shift_attention_mask = attention_mask[..., 1:]
1260
+ shift_logits = logits[..., :-1, :][
1261
+ shift_attention_mask.to(logits.device) != 0
1262
+ ].contiguous()
1263
+ shift_labels = labels[..., 1:][
1264
+ shift_attention_mask.to(labels.device) != 0
1265
+ ].contiguous()
1266
+ else:
1267
+ shift_logits = logits[..., :-1, :].contiguous()
1268
+ shift_labels = labels[..., 1:].contiguous()
1269
+ # Flatten the tokens
1270
+ loss_fct = nn.CrossEntropyLoss()
1271
+ loss = loss_fct(
1272
+ shift_logits.view(-1, shift_logits.size(-1)),
1273
+ shift_labels.view(-1).to(shift_logits.device),
1274
+ )
1275
+
1276
+ if not return_dict:
1277
+ output = (logits,) + outputs[1:]
1278
+ return (loss,) + output if loss is not None else output
1279
+
1280
+ return LlavaCausalLMOutputWithPast(
1281
+ loss=loss,
1282
+ logits=logits,
1283
+ past_key_values=outputs.past_key_values,
1284
+ hidden_states=outputs.hidden_states,
1285
+ attentions=outputs.attentions,
1286
+ )
1287
+
1288
+ def prepare_inputs_for_generation(
1289
+ self,
1290
+ input_ids,
1291
+ past_key_values=None,
1292
+ inputs_embeds=None,
1293
+ pixel_values=None,
1294
+ grid_thws=None,
1295
+ attention_mask=None,
1296
+ **kwargs,
1297
+ ):
1298
+ if past_key_values is not None:
1299
+ if hasattr(past_key_values, "get_seq_length"):
1300
+ cache_length = past_key_values.get_seq_length()
1301
+ past_length = getattr(past_key_values, "seen_tokens", cache_length)
1302
+ else:
1303
+ cache_length = past_length = past_key_values[0][0].shape[2]
1304
+
1305
+ # Keep only the unprocessed tokens:
1306
+ # 1 - If the length of the attention_mask exceeds the length of input_ids, then we are in a setting where
1307
+ # some of the inputs are exclusively passed as part of the cache (e.g. when passing input_embeds as
1308
+ # input)
1309
+ if (
1310
+ attention_mask is not None
1311
+ and attention_mask.shape[1] > input_ids.shape[1]
1312
+ ):
1313
+ input_ids = input_ids[:, -(attention_mask.shape[1] - past_length) :]
1314
+ # 2 - If the past_length is smaller than input_ids', then input_ids holds all input tokens. We can discard
1315
+ # input_ids based on the past_length.
1316
+ elif past_length < input_ids.shape[1]:
1317
+ input_ids = input_ids[:, past_length:]
1318
+ # 3 - Otherwise (past_length >= input_ids.shape[1]), let's assume input_ids only has unprocessed tokens.
1319
+ elif self.config.media_placeholder_token_id in input_ids:
1320
+ input_ids = input_ids[:, input_ids.shape[1] - 1 :]
1321
+ # If the cache has seen more tokens than it can hold, then the cache has a size limit. Let's discard the
1322
+ # older attention values, as their corresponding values are not part of the input.
1323
+ if cache_length < past_length and attention_mask is not None:
1324
+ attention_mask = attention_mask[
1325
+ :, -(cache_length + input_ids.shape[1]) :
1326
+ ]
1327
+
1328
+ position_ids = kwargs.get("position_ids", None)
1329
+ if attention_mask is not None and position_ids is None:
1330
+ # create position_ids on the fly for batch generation
1331
+ position_ids = attention_mask.long().cumsum(-1) - 1
1332
+ position_ids.masked_fill_(attention_mask == 0, 1)
1333
+ if past_key_values:
1334
+ position_ids = position_ids[:, -input_ids.shape[1] :]
1335
+
1336
+ # if `inputs_embeds` are passed, we only want to use them in the 1st generation step
1337
+ if inputs_embeds is not None and past_key_values is None:
1338
+ model_inputs = {"inputs_embeds": inputs_embeds}
1339
+ else:
1340
+ model_inputs = {"input_ids": input_ids}
1341
+
1342
+ model_inputs.update(
1343
+ {
1344
+ "position_ids": position_ids,
1345
+ "past_key_values": past_key_values,
1346
+ "use_cache": kwargs.get("use_cache"),
1347
+ "attention_mask": attention_mask,
1348
+ "pixel_values": pixel_values,
1349
+ "grid_thws": grid_thws,
1350
+ }
1351
+ )
1352
+ return model_inputs
1353
+
1354
+ def _reorder_cache(self, *args, **kwargs):
1355
+ return self.language_model._reorder_cache(*args, **kwargs)
modeling_kimi_k3_linear.py ADDED
@@ -0,0 +1,1459 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # coding=utf-8
2
+ # Copyright 2025-2026 The Moonshot AI Team, DeepSeek-AI, and HuggingFace Inc. team. All rights reserved.
3
+ #
4
+ # The multi-head latent attention, MoE gating and sparse MoE block in this file are
5
+ # adapted from DeepSeek-V3 (DeepSeek-V3/modeling_deepseek.py). They have been
6
+ # extensively modified and extended for the Kimi-Linear architecture.
7
+ #
8
+ # Licensing Information:
9
+ # - Code adapted from DeepSeek-V3 (DeepSeek-V3/modeling_deepseek.py) is licensed under the Apache License, Version 2.0.
10
+ # - Other parts of the code are licensed under the Kimi K3 License (see the LICENSE file in this repository).
11
+ #
12
+ # Apache License, Version 2.0:
13
+ # Licensed under the Apache License, Version 2.0 (the "License");
14
+ # you may not use this file except in compliance with the License.
15
+ # You may obtain a copy of the License at
16
+ #
17
+ # http://www.apache.org/licenses/LICENSE-2.0
18
+ #
19
+ # Unless required by applicable law or agreed to in writing, software
20
+ # distributed under the License is distributed on an "AS IS" BASIS,
21
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
22
+ # See the License for the specific language governing permissions and
23
+ # limitations under the License.
24
+ import math
25
+ from collections.abc import Callable
26
+ from typing import Any
27
+
28
+ import torch
29
+ import torch.nn.functional as F
30
+ import transformers
31
+ from einops import rearrange
32
+ from packaging import version
33
+ from torch import nn
34
+ from transformers.activations import ACT2FN
35
+ from transformers.cache_utils import Cache
36
+ from transformers.generation import GenerationMixin
37
+ from transformers.masking_utils import create_causal_mask
38
+ from transformers.modeling_flash_attention_utils import FlashAttentionKwargs
39
+ from transformers.modeling_outputs import (
40
+ BaseModelOutputWithPast,
41
+ CausalLMOutputWithPast,
42
+ )
43
+ from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel
44
+ from transformers.processing_utils import Unpack
45
+ from transformers.pytorch_utils import ALL_LAYERNORM_LAYERS
46
+ from transformers.utils import (
47
+ TransformersKwargs,
48
+ can_return_tuple,
49
+ logging,
50
+ )
51
+ from transformers.utils.generic import check_model_inputs
52
+ from transformers.utils.output_capturing import OutputRecorder
53
+
54
+ try:
55
+ from fla.modules import FusedRMSNormGated, ShortConvolution
56
+ from fla.ops.kda import chunk_kda, fused_recurrent_kda
57
+
58
+ # from fla.ops.kda.gate import fused_kda_gate # deprecated, gate is now computed inside chunk_kda/fused_recurrent_kda
59
+ from fla.ops.utils.index import prepare_cu_seqlens_from_mask, prepare_lens_from_mask
60
+ from fla.utils import tensor_cache
61
+ except ImportError:
62
+ raise ImportError("Plese run `pip install -U fla-core`")
63
+
64
+ def get_calibrate_all_experts_flag() -> bool:
65
+ return False
66
+
67
+ from .configuration_kimi_k3 import KimiLinearConfig
68
+
69
+ assert version.parse(transformers.__version__) >= version.parse(
70
+ "4.56.0"
71
+ ), "Please upgrade transformers to >= 4.56.0"
72
+
73
+ logger = logging.get_logger(__name__)
74
+
75
+
76
+ # Register Moonshot-specific activation functions
77
+ class SituAndMul(nn.Module):
78
+ """
79
+ SituAndMul activation: beta * tanh(gate / beta) * sigmoid(gate) * up
80
+ When linear_beta is set, up is also transformed by linear_beta * tanh(up / linear_beta).
81
+ """
82
+
83
+ def __init__(self, beta: float = 1.0, linear_beta: float | None = None):
84
+ super().__init__()
85
+ self.beta = beta
86
+ self.linear_beta = linear_beta
87
+
88
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
89
+ d = x.shape[-1] // 2
90
+ gate = x[..., :d].to(torch.float32)
91
+ up = x[..., d:].to(torch.float32)
92
+ situ_a = self.beta * torch.tanh(gate / self.beta) * torch.sigmoid(gate)
93
+ if self.linear_beta is not None:
94
+ up = self.linear_beta * torch.tanh(up / self.linear_beta)
95
+ return (situ_a * up).to(x.dtype)
96
+
97
+
98
+ ACT2FN["situ"] = SituAndMul
99
+
100
+
101
+ def _get_situ_activation_params(config: KimiLinearConfig):
102
+ beta = getattr(config, "activation_situ_beta", None)
103
+ linear_beta = getattr(config, "activation_situ_linear_beta", None)
104
+ return beta or 1.0, linear_beta
105
+
106
+
107
+ def index_first_axis(x, indices):
108
+ return x[indices]
109
+
110
+
111
+ @tensor_cache
112
+ def get_unpad_data(
113
+ attention_mask: torch.Tensor,
114
+ ) -> tuple[torch.Tensor, torch.Tensor, int]:
115
+ lens = prepare_lens_from_mask(attention_mask)
116
+ indices = torch.nonzero(attention_mask.flatten(), as_tuple=False).flatten()
117
+ max_seqlen_in_batch = lens.max().item()
118
+ cu_seqlens = prepare_cu_seqlens_from_mask(attention_mask)
119
+ return indices, cu_seqlens, max_seqlen_in_batch
120
+
121
+
122
+ def pad_input(
123
+ hidden_states: torch.Tensor,
124
+ indices: torch.LongTensor,
125
+ batch_size: int,
126
+ seq_len: int,
127
+ ) -> torch.Tensor:
128
+ out = hidden_states.new_zeros((batch_size * seq_len, *hidden_states.shape[1:]))
129
+ out[indices] = hidden_states
130
+ return out.view(batch_size, seq_len, *hidden_states.shape[1:])
131
+
132
+
133
+ class KimiDynamicCache:
134
+ """
135
+ Dynamic cache for Kimi model.
136
+ Inspired by Qwen3-Next
137
+ """
138
+
139
+ is_compileable = False
140
+
141
+ def __init__(self, config: KimiLinearConfig):
142
+ super().__init__()
143
+ self.config = config
144
+
145
+ if config.linear_attn_config is not None:
146
+ self.layer_types = []
147
+ for i in range(config.num_hidden_layers):
148
+ if config.is_kda_layer(i):
149
+ self.layer_types.append("linear_attention")
150
+ else:
151
+ self.layer_types.append("full_attention")
152
+ else:
153
+ self.layer_types = ["full_attention"] * config.num_hidden_layers
154
+
155
+ self.transformer_layers = [
156
+ i
157
+ for i in range(config.num_hidden_layers)
158
+ if self.layer_types[i] == "full_attention"
159
+ ]
160
+
161
+ linear_layers = [
162
+ i
163
+ for i in range(config.num_hidden_layers)
164
+ if self.layer_types[i] == "linear_attention"
165
+ ]
166
+ self.last_linear_layer = linear_layers[-1] if linear_layers else -1
167
+
168
+ self.conv_states = [None for _ in range(config.num_hidden_layers)]
169
+ self.recurrent_states = [None for _ in range(config.num_hidden_layers)]
170
+ self.key_cache = [None for _ in range(config.num_hidden_layers)]
171
+ self.value_cache = [None for _ in range(config.num_hidden_layers)]
172
+
173
+ def __len__(self):
174
+ return len(self.layer_types)
175
+
176
+ def update(
177
+ self,
178
+ key_states: torch.Tensor,
179
+ value_states: torch.Tensor,
180
+ layer_idx: int,
181
+ cache_kwargs: dict[str, Any] | None = None,
182
+ ) -> tuple[torch.Tensor, torch.Tensor]:
183
+ if self.key_cache[layer_idx] is None:
184
+ self.key_cache[layer_idx] = key_states
185
+ self.value_cache[layer_idx] = value_states
186
+ else:
187
+ self.key_cache[layer_idx] = torch.cat(
188
+ [self.key_cache[layer_idx], key_states], dim=2
189
+ )
190
+ self.value_cache[layer_idx] = torch.cat(
191
+ [self.value_cache[layer_idx], value_states], dim=2
192
+ )
193
+
194
+ return self.key_cache[layer_idx], self.value_cache[layer_idx]
195
+
196
+ def reorder_cache(self, beam_idx: torch.LongTensor):
197
+ """Reorders the cache for beam search, given the selected beam indices."""
198
+ for layer_idx in range(len(self.key_cache)):
199
+ if self.key_cache[layer_idx] is not None:
200
+ device = self.key_cache[layer_idx].device
201
+ beam_idx = beam_idx.to(device)
202
+ self.key_cache[layer_idx] = self.key_cache[layer_idx].index_select(
203
+ 0, beam_idx
204
+ )
205
+ self.value_cache[layer_idx] = self.value_cache[layer_idx].index_select(
206
+ 0, beam_idx
207
+ )
208
+
209
+ if self.conv_states[layer_idx] is not None:
210
+ device = self.conv_states[layer_idx][0].device
211
+ beam_idx = beam_idx.to(device)
212
+ q_conv, k_conv, v_conv = self.conv_states[layer_idx]
213
+ self.conv_states[layer_idx] = (
214
+ q_conv.index_select(0, beam_idx),
215
+ k_conv.index_select(0, beam_idx),
216
+ v_conv.index_select(0, beam_idx),
217
+ )
218
+ self.recurrent_states[layer_idx] = self.recurrent_states[
219
+ layer_idx
220
+ ].index_select(0, beam_idx)
221
+
222
+ def get_seq_length(self, layer_idx: int | None = 0) -> int:
223
+ """Returns the sequence length of the cached states. A layer index can be optionally passed."""
224
+ # take any layer that contains cache and not empty tensor
225
+ layer_idx = (
226
+ self.transformer_layers[0]
227
+ if layer_idx not in self.transformer_layers
228
+ else layer_idx
229
+ )
230
+ if len(self.key_cache) <= layer_idx or self.key_cache[layer_idx] is None:
231
+ return 0
232
+ return self.key_cache[layer_idx].shape[-2]
233
+
234
+ def get_query_offset(self, layer_idx: int = 0) -> int:
235
+ return self.get_seq_length(layer_idx=layer_idx)
236
+
237
+ def get_mask_sizes(self, cache_position, layer_idx: int) -> tuple[int, int]:
238
+ """
239
+ Return a tuple (kv_length, kv_offset) corresponding to the length and offset that will be returned for
240
+ the given layer at `layer_idx`.
241
+ The masks are then prepared according to the given lengths (kv_length, kv_offset) and patterns for each layer.
242
+ """
243
+ kv_offset = 0
244
+ # cache_position may be an int (new API) or a 1-D tensor (old API)
245
+ query_length = (
246
+ cache_position
247
+ if isinstance(cache_position, int)
248
+ else cache_position.shape[0]
249
+ )
250
+ past_seen_tokens = self.get_seq_length(layer_idx)
251
+ kv_length = query_length + past_seen_tokens
252
+ return kv_length, kv_offset
253
+
254
+ @property
255
+ def has_previous_state(self):
256
+ """We have a previous state if the last linear (conv) layer was already updated."""
257
+ if self.last_linear_layer == -1:
258
+ return False
259
+ return self.conv_states[self.last_linear_layer] is not None
260
+
261
+
262
+ class KimiRMSNorm(nn.Module):
263
+ def __init__(self, hidden_size, eps=1e-6):
264
+ super().__init__()
265
+ self.weight = nn.Parameter(torch.ones(hidden_size))
266
+ self.variance_epsilon = eps
267
+
268
+ def forward(self, hidden_states):
269
+ dtype = hidden_states.dtype
270
+ x = hidden_states.float()
271
+ x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.variance_epsilon)
272
+ return self.weight * x.to(dtype)
273
+
274
+
275
+ ALL_LAYERNORM_LAYERS.append(KimiRMSNorm)
276
+
277
+
278
+ class KimiBlockSparseMLP(nn.Module):
279
+ def __init__(
280
+ self, config: KimiLinearConfig, hidden_size=None, intermediate_size=None
281
+ ):
282
+ super().__init__()
283
+ self.config = config
284
+ self.ffn_dim = (
285
+ config.intermediate_size if intermediate_size is None else intermediate_size
286
+ )
287
+ self.hidden_dim = config.hidden_size if hidden_size is None else hidden_size
288
+
289
+ self.w1 = nn.Linear(self.hidden_dim, self.ffn_dim, bias=False) # gate
290
+ self.w2 = nn.Linear(self.ffn_dim, self.hidden_dim, bias=False) # down
291
+ self.w3 = nn.Linear(self.hidden_dim, self.ffn_dim, bias=False) # up
292
+
293
+ if config.hidden_act == "situ":
294
+ beta, linear_beta = _get_situ_activation_params(config)
295
+ self.act_fn = SituAndMul(
296
+ beta=beta,
297
+ linear_beta=linear_beta,
298
+ )
299
+ else:
300
+ self.act_fn = ACT2FN[config.hidden_act]
301
+
302
+ def forward(self, hidden_states):
303
+ if self.config.hidden_act == "situ":
304
+ gate_up = torch.cat(
305
+ [self.w1(hidden_states), self.w3(hidden_states)], dim=-1
306
+ )
307
+ current_hidden_states = self.act_fn(gate_up)
308
+ else:
309
+ current_hidden_states = self.act_fn(self.w1(hidden_states)) * self.w3(
310
+ hidden_states
311
+ )
312
+ current_hidden_states = self.w2(current_hidden_states)
313
+ return current_hidden_states
314
+
315
+
316
+ class KimiMLP(nn.Module):
317
+ def __init__(
318
+ self, config: KimiLinearConfig, hidden_size=None, intermediate_size=None
319
+ ):
320
+ super().__init__()
321
+ self.config = config
322
+ self.hidden_size = config.hidden_size if hidden_size is None else hidden_size
323
+ self.intermediate_size = (
324
+ config.intermediate_size if intermediate_size is None else intermediate_size
325
+ )
326
+ self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
327
+ self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
328
+ self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)
329
+ if config.hidden_act == "situ":
330
+ beta, linear_beta = _get_situ_activation_params(config)
331
+ self.act_fn = SituAndMul(
332
+ beta=beta,
333
+ linear_beta=linear_beta,
334
+ )
335
+ else:
336
+ self.act_fn = ACT2FN[config.hidden_act]
337
+
338
+ def forward(self, x):
339
+ if self.config.hidden_act == "situ":
340
+ gate_up = torch.cat([self.gate_proj(x), self.up_proj(x)], dim=-1)
341
+ down_proj = self.down_proj(self.act_fn(gate_up))
342
+ else:
343
+ down_proj = self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
344
+ return down_proj
345
+
346
+
347
+ def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
348
+ """Expand the key/value heads from `num_key_value_heads` to `num_attention_heads`."""
349
+ if n_rep == 1:
350
+ return hidden_states
351
+ return torch.repeat_interleave(hidden_states, dim=1, repeats=n_rep)
352
+
353
+
354
+ def eager_attention_forward(
355
+ module: nn.Module,
356
+ query: torch.Tensor,
357
+ key: torch.Tensor,
358
+ value: torch.Tensor,
359
+ attention_mask: torch.Tensor | None,
360
+ scaling: float,
361
+ dropout: float = 0.0,
362
+ **kwargs: Unpack[TransformersKwargs],
363
+ ):
364
+ key = repeat_kv(key, module.num_key_value_groups)
365
+ value = repeat_kv(value, module.num_key_value_groups)
366
+
367
+ scores = torch.einsum("bhqd,bhkd->bhqk", query, key) * scaling
368
+ if attention_mask is not None:
369
+ scores = scores + attention_mask[:, :, :, : key.shape[-2]]
370
+
371
+ probs = F.softmax(scores, dim=-1, dtype=torch.float32).to(query.dtype)
372
+ probs = F.dropout(probs, p=dropout, training=module.training)
373
+ out = torch.einsum("bhqk,bhkd->bhqd", probs, value).transpose(1, 2).contiguous()
374
+
375
+ return out, probs
376
+
377
+
378
+ class KimiMLAAttention(nn.Module):
379
+ """
380
+ Multi-Latent Attention adapted from deepseek-v3
381
+ """
382
+
383
+ def __init__(self, config: KimiLinearConfig, layer_idx: int):
384
+ nn.Module.__init__(self)
385
+ self.config = config
386
+ self.layer_idx = layer_idx
387
+ self.hidden_size = config.hidden_size
388
+ self.num_heads = config.num_attention_heads
389
+ self.num_key_value_heads = config.num_key_value_heads
390
+ self.num_key_value_groups = self.num_heads // self.num_key_value_heads
391
+
392
+ self.attention_dropout = getattr(config, "attention_dropout", 0.0)
393
+
394
+ try:
395
+ self.q_lora_rank = config.q_lora_rank
396
+ self.qk_rope_head_dim = config.qk_rope_head_dim
397
+ self.kv_lora_rank = config.kv_lora_rank
398
+ self.v_head_dim = config.v_head_dim
399
+ self.qk_nope_head_dim = config.qk_nope_head_dim
400
+ self.q_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim
401
+ self.use_nope = config.mla_use_nope
402
+ self.scaling = self.q_head_dim ** (-0.5)
403
+ except Exception as e:
404
+ raise ValueError(
405
+ f"Kimi MLA config is not found or not properly formatted: {e}"
406
+ )
407
+
408
+ if self.q_lora_rank is not None:
409
+ self.q_a_proj = nn.Linear(
410
+ self.hidden_size,
411
+ self.q_lora_rank,
412
+ bias=False,
413
+ )
414
+ self.q_a_layernorm = KimiRMSNorm(self.q_lora_rank)
415
+ self.q_b_proj = nn.Linear(
416
+ self.q_lora_rank,
417
+ self.num_heads * self.q_head_dim,
418
+ bias=False,
419
+ )
420
+ else:
421
+ self.q_proj = nn.Linear(
422
+ self.hidden_size,
423
+ self.num_heads * self.q_head_dim,
424
+ bias=False,
425
+ )
426
+ self.kv_a_proj_with_mqa = nn.Linear(
427
+ self.hidden_size,
428
+ self.kv_lora_rank + self.qk_rope_head_dim,
429
+ bias=False,
430
+ )
431
+ self.kv_a_layernorm = KimiRMSNorm(self.kv_lora_rank)
432
+ self.kv_b_proj = nn.Linear(
433
+ self.kv_lora_rank,
434
+ self.num_heads
435
+ * (self.q_head_dim - self.qk_rope_head_dim + self.v_head_dim),
436
+ bias=False,
437
+ )
438
+ self.o_proj = nn.Linear(
439
+ self.num_heads * self.v_head_dim,
440
+ self.hidden_size,
441
+ bias=False,
442
+ )
443
+ self.is_causal = True
444
+ assert self.use_nope
445
+
446
+ self.use_output_gate = getattr(config, "mla_use_output_gate", False)
447
+ if self.use_output_gate:
448
+ projection_size = self.num_heads * self.v_head_dim
449
+ self.g_proj = nn.Linear(self.hidden_size, projection_size, bias=False)
450
+
451
+ self.rotary_emb = None
452
+
453
+ def forward(
454
+ self,
455
+ hidden_states: torch.Tensor,
456
+ attention_mask: torch.Tensor | None = None,
457
+ position_ids: torch.LongTensor | None = None,
458
+ past_key_values: Cache | None = None,
459
+ **kwargs,
460
+ ) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]:
461
+ batch_size, seq_length = hidden_states.shape[:-1]
462
+ query_shape = (batch_size, seq_length, -1, self.q_head_dim)
463
+ key_shape = (
464
+ batch_size,
465
+ seq_length,
466
+ -1,
467
+ self.qk_nope_head_dim + self.v_head_dim,
468
+ )
469
+
470
+ if self.q_lora_rank is not None:
471
+ q_states = self.q_b_proj(self.q_a_layernorm(self.q_a_proj(hidden_states)))
472
+ else:
473
+ q_states = self.q_proj(hidden_states)
474
+ q_states = q_states.view(query_shape).transpose(1, 2)
475
+ q_pass, q_rot = torch.split(
476
+ q_states, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1
477
+ )
478
+
479
+ compressed_kv = self.kv_a_proj_with_mqa(hidden_states)
480
+ k_pass, k_rot = torch.split(
481
+ compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1
482
+ )
483
+
484
+ k_pass = (
485
+ self.kv_b_proj(self.kv_a_layernorm(k_pass)).view(key_shape).transpose(1, 2)
486
+ )
487
+ k_pass, value_states = torch.split(
488
+ k_pass, [self.qk_nope_head_dim, self.v_head_dim], dim=-1
489
+ )
490
+
491
+ k_rot = k_rot.view(batch_size, 1, seq_length, self.qk_rope_head_dim)
492
+
493
+ k_rot = k_rot.expand(*k_pass.shape[:-1], -1)
494
+
495
+ query_states = torch.cat((q_pass, q_rot), dim=-1)
496
+ key_states = torch.cat((k_pass, k_rot), dim=-1)
497
+
498
+ if past_key_values is not None:
499
+ key_states, value_states = past_key_values.update(
500
+ key_states, value_states, self.layer_idx
501
+ )
502
+
503
+ if (
504
+ self.config._attn_implementation == "flash_attention_2"
505
+ and self.q_head_dim != self.v_head_dim
506
+ ):
507
+ value_states = F.pad(value_states, [0, self.q_head_dim - self.v_head_dim])
508
+
509
+ attention_interface: Callable = eager_attention_forward
510
+ if self.config._attn_implementation != "eager":
511
+ attention_interface = ALL_ATTENTION_FUNCTIONS[
512
+ self.config._attn_implementation
513
+ ]
514
+
515
+ attn_output, _ = attention_interface(
516
+ self,
517
+ query_states,
518
+ key_states,
519
+ value_states,
520
+ attention_mask,
521
+ dropout=0.0 if not self.training else self.attention_dropout,
522
+ scaling=self.scaling,
523
+ **kwargs,
524
+ )
525
+
526
+ if (
527
+ self.config._attn_implementation == "flash_attention_2"
528
+ and self.q_head_dim != self.v_head_dim
529
+ ):
530
+ attn_output = attn_output[:, :, :, : self.v_head_dim]
531
+
532
+ attn_output = attn_output.reshape(batch_size, seq_length, -1).contiguous()
533
+ if self.use_output_gate:
534
+ g = self.g_proj(hidden_states).sigmoid()
535
+ attn_output = attn_output * g
536
+ attn_output = self.o_proj(attn_output)
537
+ return attn_output
538
+
539
+
540
+ class KimiDeltaAttention(nn.Module):
541
+ def __init__(self, config: KimiLinearConfig, layer_idx: int):
542
+ super().__init__()
543
+ self.config = config
544
+ self.mode = "chunk"
545
+
546
+ self.hidden_size = config.hidden_size
547
+ self.conv_size = config.linear_attn_config["short_conv_kernel_size"]
548
+ self.head_dim = config.linear_attn_config["head_dim"]
549
+ self.num_heads = config.linear_attn_config["num_heads"]
550
+ self.head_k_dim = self.head_dim
551
+ self.num_k_heads = self.num_heads
552
+
553
+ self.layer_idx = layer_idx
554
+
555
+ assert self.mode in [
556
+ "chunk",
557
+ "fused_recurrent",
558
+ ], f"Not supported mode `{self.mode}`."
559
+
560
+ projection_k_size = self.head_k_dim * self.num_k_heads
561
+ projection_size = self.head_dim * self.num_heads
562
+
563
+ self.q_proj = nn.Linear(self.hidden_size, projection_k_size, bias=False)
564
+ self.k_proj = nn.Linear(self.hidden_size, projection_k_size, bias=False)
565
+ self.v_proj = nn.Linear(self.hidden_size, projection_size, bias=False)
566
+
567
+ self.q_conv1d = ShortConvolution(
568
+ hidden_size=projection_k_size,
569
+ kernel_size=self.conv_size,
570
+ activation="silu",
571
+ )
572
+ self.k_conv1d = ShortConvolution(
573
+ hidden_size=projection_k_size,
574
+ kernel_size=self.conv_size,
575
+ activation="silu",
576
+ )
577
+ self.v_conv1d = ShortConvolution(
578
+ hidden_size=projection_size,
579
+ kernel_size=self.conv_size,
580
+ activation="silu",
581
+ )
582
+
583
+ self.A_log = torch.nn.Parameter(
584
+ torch.log(torch.empty(self.num_heads, dtype=torch.float32).uniform_(1, 16))
585
+ )
586
+
587
+ self.f_a_proj = nn.Linear(self.hidden_size, self.head_dim, bias=False)
588
+ self.f_b_proj = nn.Linear(self.head_dim, projection_size, bias=False)
589
+
590
+ self.dt_bias = nn.Parameter(torch.empty(projection_size, dtype=torch.float32))
591
+
592
+ self.b_proj = nn.Linear(self.hidden_size, self.num_heads, bias=False)
593
+
594
+ self.use_full_rank_gate = config.linear_attn_config.get(
595
+ "use_full_rank_gate", False
596
+ )
597
+ self.gate_lower_bound = config.linear_attn_config.get("gate_lower_bound", None)
598
+ if self.use_full_rank_gate:
599
+ self.g_proj = nn.Linear(self.hidden_size, projection_size, bias=False)
600
+ else:
601
+ self.g_a_proj = nn.Linear(self.hidden_size, self.head_dim, bias=False)
602
+ self.g_b_proj = nn.Linear(self.head_dim, projection_size, bias=False)
603
+
604
+ self.o_norm = FusedRMSNormGated(
605
+ self.head_dim, eps=config.rms_norm_eps, activation="sigmoid"
606
+ )
607
+ self.o_proj = nn.Linear(projection_size, self.hidden_size, bias=False)
608
+
609
+ def forward(
610
+ self,
611
+ hidden_states: torch.Tensor,
612
+ attention_mask: torch.Tensor | None = None,
613
+ cache_params: KimiDynamicCache | None = None,
614
+ **kwargs: Unpack[dict],
615
+ ) -> tuple[torch.Tensor, torch.Tensor | None, Cache | None]:
616
+ if attention_mask is not None:
617
+ if attention_mask.dim() != 2:
618
+ attention_mask = kwargs.get("padding_mask")
619
+
620
+ if attention_mask is not None and attention_mask.dim() != 2:
621
+ raise ValueError(
622
+ "attention_mask must be a 0-1 matrix of shape [batch_size, seq_len] "
623
+ "(0 = padding). 3D masks are not supported here.",
624
+ )
625
+ use_cache = cache_params is not None
626
+ batch_size, q_len, _ = hidden_states.shape
627
+ mode = "fused_recurrent" if use_cache and q_len == 1 else self.mode
628
+ if self.training:
629
+ assert mode == "chunk", "Only chunk mode is supported in training."
630
+
631
+ cu_seqlens = kwargs.get("cu_seqlens")
632
+ indices = None
633
+ if attention_mask is not None:
634
+ indices, cu_seqlens, _ = get_unpad_data(attention_mask[:, -q_len:])
635
+ hidden_states = index_first_axis(
636
+ rearrange(hidden_states, "b s ... -> (b s) ..."), indices
637
+ ).unsqueeze(0)
638
+
639
+ conv_state_q, conv_state_k, conv_state_v = None, None, None
640
+ recurrent_state = None
641
+ if cache_params is not None:
642
+ if cache_params.conv_states[self.layer_idx] is not None:
643
+ conv_state_q, conv_state_k, conv_state_v = cache_params.conv_states[
644
+ self.layer_idx
645
+ ]
646
+ recurrent_state = cache_params.recurrent_states[self.layer_idx]
647
+
648
+ q_proj_states = self.q_proj(hidden_states)
649
+ k_proj_states = self.k_proj(hidden_states)
650
+ v_proj_states = self.v_proj(hidden_states)
651
+ q, conv_state_q = self.q_conv1d(
652
+ x=q_proj_states,
653
+ cache=conv_state_q,
654
+ output_final_state=use_cache,
655
+ cu_seqlens=cu_seqlens,
656
+ )
657
+ k, conv_state_k = self.k_conv1d(
658
+ x=k_proj_states,
659
+ cache=conv_state_k,
660
+ output_final_state=use_cache,
661
+ cu_seqlens=cu_seqlens,
662
+ )
663
+ v, conv_state_v = self.v_conv1d(
664
+ x=v_proj_states,
665
+ cache=conv_state_v,
666
+ output_final_state=use_cache,
667
+ cu_seqlens=cu_seqlens,
668
+ )
669
+ g = self.f_b_proj(self.f_a_proj(hidden_states))
670
+ g = rearrange(g, "... (h d) -> ... h d", d=self.head_dim)
671
+ beta = self.b_proj(hidden_states).float()
672
+
673
+ q, k = map(
674
+ lambda x: rearrange(x, "... (h d) -> ... h d", d=self.head_k_dim), (q, k)
675
+ )
676
+ v = rearrange(v, "... (h d) -> ... h d", d=self.head_dim)
677
+
678
+ if mode == "chunk":
679
+ o, recurrent_state = chunk_kda(
680
+ q=q,
681
+ k=k,
682
+ v=v,
683
+ g=g,
684
+ beta=beta,
685
+ A_log=self.A_log,
686
+ dt_bias=self.dt_bias,
687
+ initial_state=recurrent_state,
688
+ output_final_state=True,
689
+ use_qk_l2norm_in_kernel=True,
690
+ use_gate_in_kernel=True,
691
+ use_beta_sigmoid_in_kernel=True,
692
+ safe_gate=self.gate_lower_bound is not None,
693
+ lower_bound=self.gate_lower_bound,
694
+ transpose_state_layout=True,
695
+ cu_seqlens=cu_seqlens,
696
+ )
697
+ else:
698
+ o, recurrent_state = fused_recurrent_kda(
699
+ q=q,
700
+ k=k,
701
+ v=v,
702
+ g=g,
703
+ beta=beta,
704
+ A_log=self.A_log,
705
+ dt_bias=self.dt_bias,
706
+ initial_state=recurrent_state,
707
+ output_final_state=True,
708
+ use_qk_l2norm_in_kernel=True,
709
+ use_gate_in_kernel=True,
710
+ use_beta_sigmoid_in_kernel=True,
711
+ lower_bound=self.gate_lower_bound,
712
+ transpose_state_layout=True,
713
+ cu_seqlens=cu_seqlens,
714
+ )
715
+ if cache_params is not None:
716
+ cache_params.recurrent_states[self.layer_idx] = recurrent_state
717
+ cache_params.conv_states[self.layer_idx] = (
718
+ conv_state_q,
719
+ conv_state_k,
720
+ conv_state_v,
721
+ )
722
+
723
+ if self.use_full_rank_gate:
724
+ g = self.g_proj(hidden_states)
725
+ else:
726
+ g = self.g_b_proj(self.g_a_proj(hidden_states))
727
+ g = rearrange(g, "... (h d) -> ... h d", d=self.head_dim)
728
+ o = self.o_norm(o, g)
729
+
730
+ o = rearrange(o, "b t h d -> b t (h d)")
731
+ o = self.o_proj(o)
732
+ if attention_mask is not None:
733
+ o = pad_input(o.squeeze(0), indices, batch_size, q_len)
734
+
735
+ return o
736
+
737
+
738
+ class KimiMoEGate(nn.Module):
739
+ """
740
+ MoEGate adapted from Deepseek-V3.
741
+ Parameter correspondences:
742
+ num_experts -> n_routed_experts
743
+ num_experts_per_token -> num_experts_per_tok
744
+ num_expert_group -> n_group
745
+ moe_router_activation_func -> scoring_func
746
+ """
747
+
748
+ def __init__(self, config: KimiLinearConfig):
749
+ super().__init__()
750
+ self.config = config
751
+ self.top_k = config.num_experts_per_token
752
+ self.num_experts = config.num_experts
753
+ self.routed_scaling_factor = config.routed_scaling_factor
754
+ self.moe_router_activation_func = config.moe_router_activation_func
755
+ self.num_expert_group = getattr(config, "num_expert_group", 1)
756
+ self.topk_group = getattr(config, "topk_group", 1)
757
+
758
+ # topk selection algorithm
759
+ self.moe_renormalize = config.moe_renormalize
760
+ self.gating_dim = config.hidden_size
761
+ self.weight = nn.Parameter(
762
+ torch.empty((self.num_experts, self.gating_dim)),
763
+ )
764
+
765
+ self.e_score_correction_bias = nn.Parameter(
766
+ torch.empty(self.num_experts),
767
+ )
768
+ self.reset_parameters()
769
+
770
+ def reset_parameters(self) -> None:
771
+ import torch.nn.init as init
772
+
773
+ init.kaiming_uniform_(self.weight, a=math.sqrt(5))
774
+
775
+ def forward(self, hidden_states):
776
+ bsz, seq_len, h = hidden_states.shape
777
+ # compute gating score
778
+ hidden_states = hidden_states.view(-1, h)
779
+ logits = F.linear(
780
+ hidden_states.type(torch.float32),
781
+ self.weight.type(torch.float32),
782
+ None,
783
+ )
784
+ if self.moe_router_activation_func == "sigmoid":
785
+ scores = logits.sigmoid()
786
+ elif self.moe_router_activation_func == "softmax":
787
+ scores = logits.softmax(dim=1)
788
+ else:
789
+ raise NotImplementedError(
790
+ f"insupportable scoring function for MoE gating: {self.moe_router_activation_func}",
791
+ )
792
+
793
+ # select top-k experts
794
+ scores = scores.view(bsz * seq_len, -1)
795
+ scores_for_choice = scores + self.e_score_correction_bias.unsqueeze(0)
796
+ if self.num_expert_group > 1 and self.num_expert_group > self.topk_group:
797
+ group_scores = (
798
+ scores_for_choice.view(bsz * seq_len, self.num_expert_group, -1)
799
+ .topk(2, dim=-1)[0]
800
+ .sum(dim=-1)
801
+ ) # [n, num_expert_group]
802
+ group_idx = torch.topk(
803
+ group_scores,
804
+ k=self.topk_group,
805
+ dim=-1,
806
+ sorted=False,
807
+ )[1] # [n, top_k_group]
808
+ group_mask = torch.zeros_like(group_scores) # [n, num_expert_group]
809
+ group_mask.scatter_(1, group_idx, 1) # [n, num_expert_group]
810
+ score_mask = (
811
+ group_mask.unsqueeze(-1)
812
+ .expand(
813
+ bsz * seq_len,
814
+ self.num_expert_group,
815
+ self.num_experts // self.num_expert_group,
816
+ )
817
+ .reshape(bsz * seq_len, -1)
818
+ ) # [n, e]
819
+ tmp_scores = scores_for_choice.masked_fill(
820
+ ~score_mask.bool(), float("-inf")
821
+ ) # [n, e]
822
+ else:
823
+ tmp_scores = scores_for_choice
824
+ _, topk_idx = torch.topk(
825
+ tmp_scores,
826
+ k=self.top_k,
827
+ dim=-1,
828
+ sorted=False,
829
+ )
830
+ topk_weight = scores.gather(1, topk_idx)
831
+
832
+ # norm gate to sum 1
833
+ if self.top_k > 1 and self.moe_renormalize:
834
+ denominator = topk_weight.sum(dim=-1, keepdim=True) + 1e-20
835
+ topk_weight = topk_weight / denominator
836
+ # must multiply the scaling factor
837
+ topk_weight = topk_weight * self.routed_scaling_factor
838
+
839
+ return topk_idx, topk_weight
840
+
841
+
842
+ class KimiSparseMoeBlock(nn.Module):
843
+ """
844
+ Adapted from Deepseek-V3's MOE implementation
845
+ The namings are consistent with Kimi's version.
846
+ """
847
+
848
+ def __init__(self, config: KimiLinearConfig):
849
+ super().__init__()
850
+ self.config = config
851
+ self.hidden_dim = config.hidden_size
852
+ self.num_experts = config.num_experts
853
+ self.top_k = config.num_experts_per_token
854
+ self.moe_renormalize = config.moe_renormalize
855
+
856
+ self.use_latent_moe = (
857
+ getattr(config, "routed_expert_hidden_size", None) is not None
858
+ )
859
+ self.moe_hidden_size = (
860
+ config.routed_expert_hidden_size
861
+ if self.use_latent_moe
862
+ else config.hidden_size
863
+ )
864
+ self.latent_moe_use_norm = getattr(config, "latent_moe_use_norm", False)
865
+
866
+ self.ep_size = 1
867
+ self.experts_per_rank = config.num_experts
868
+ self.ep_rank = 0
869
+ self.experts = nn.ModuleList(
870
+ [
871
+ KimiBlockSparseMLP(
872
+ config,
873
+ hidden_size=self.moe_hidden_size,
874
+ intermediate_size=config.moe_intermediate_size,
875
+ )
876
+ for _ in range(config.num_experts)
877
+ ],
878
+ )
879
+ self.gate = KimiMoEGate(config)
880
+ if config.num_shared_experts is not None:
881
+ intermediate_size = config.moe_intermediate_size * config.num_shared_experts
882
+ self.shared_experts = KimiMLP(
883
+ config=config,
884
+ intermediate_size=intermediate_size,
885
+ )
886
+
887
+ if self.use_latent_moe:
888
+ self.routed_expert_down_proj = nn.Linear(
889
+ config.hidden_size,
890
+ self.moe_hidden_size,
891
+ bias=False,
892
+ )
893
+ self.routed_expert_up_proj = nn.Linear(
894
+ self.moe_hidden_size,
895
+ config.hidden_size,
896
+ bias=False,
897
+ )
898
+ if self.latent_moe_use_norm:
899
+ self.routed_expert_norm = KimiRMSNorm(
900
+ self.moe_hidden_size,
901
+ eps=config.rms_norm_eps,
902
+ )
903
+
904
+ def forward(self, hidden_states):
905
+ identity = hidden_states
906
+ orig_shape = hidden_states.shape
907
+ topk_idx, topk_weight = self.gate(hidden_states)
908
+ hidden_states = hidden_states.view(-1, hidden_states.shape[-1])
909
+
910
+ if self.use_latent_moe:
911
+ hidden_states = self.routed_expert_down_proj(hidden_states)
912
+
913
+ if not self.training:
914
+ y = self.moe_infer(hidden_states, topk_idx, topk_weight)
915
+ else:
916
+ y = self.moe_train(hidden_states, topk_idx, topk_weight)
917
+
918
+ if self.use_latent_moe:
919
+ if self.latent_moe_use_norm:
920
+ y = self.routed_expert_norm(y)
921
+ y = self.routed_expert_up_proj(y)
922
+
923
+ y = y.view(*orig_shape)
924
+
925
+ if self.config.num_shared_experts is not None:
926
+ y = y + self.shared_experts(identity)
927
+ return y
928
+
929
+ def moe_train(self, x, topk_ids, topk_weight):
930
+ """Training-compatible MoE dispatch with gradient flow."""
931
+ y = torch.zeros_like(x)
932
+
933
+ with torch.no_grad():
934
+ expert_mask = F.one_hot(topk_ids, self.num_experts).permute(2, 1, 0)
935
+
936
+ for expert_idx, expert in enumerate(self.experts):
937
+ top_k_pos, token_indices = torch.where(expert_mask[expert_idx])
938
+
939
+ if get_calibrate_all_experts_flag():
940
+ expert_out = expert(x)[token_indices]
941
+ else:
942
+ expert_out = expert(x[token_indices])
943
+
944
+ expert_weights = topk_weight[token_indices, top_k_pos, None]
945
+ y.index_add_(0, token_indices, (expert_out * expert_weights).to(y.dtype))
946
+
947
+ return y
948
+
949
+ @torch.no_grad()
950
+ def moe_infer(self, x, topk_ids, topk_weight):
951
+ cnts = topk_ids.new_zeros((topk_ids.shape[0], len(self.experts)))
952
+ cnts.scatter_(1, topk_ids, 1)
953
+ tokens_per_expert = cnts.sum(dim=0)
954
+ idxs = topk_ids.view(-1).argsort()
955
+ sorted_tokens = x[idxs // topk_ids.shape[1]]
956
+
957
+ tokens_per_expert = tokens_per_expert.cpu().numpy()
958
+
959
+ outputs = []
960
+ start_idx = 0
961
+ for i, num_tokens in enumerate(tokens_per_expert):
962
+ end_idx = start_idx + num_tokens
963
+ if num_tokens == 0:
964
+ continue
965
+ expert = self.experts[i + self.ep_rank * self.experts_per_rank]
966
+ tokens_for_this_expert = sorted_tokens[start_idx:end_idx]
967
+ expert_out = expert(tokens_for_this_expert)
968
+ outputs.append(expert_out)
969
+ start_idx = end_idx
970
+
971
+ outs = torch.cat(outputs, dim=0) if len(outputs) else sorted_tokens.new_empty(0)
972
+
973
+ new_x = torch.empty_like(outs)
974
+ new_x[idxs] = outs
975
+ final_out = (
976
+ new_x.view(*topk_ids.shape, -1)
977
+ .type(topk_weight.dtype)
978
+ .mul_(topk_weight.unsqueeze(dim=-1))
979
+ .sum(dim=1)
980
+ .type(new_x.dtype)
981
+ )
982
+ return final_out
983
+
984
+
985
+ class KimiDecoderLayer(nn.Module):
986
+ def __init__(self, config: KimiLinearConfig, layer_idx: int):
987
+ super().__init__()
988
+ self.hidden_size = config.hidden_size
989
+ self.config = config
990
+ self.layer_idx = layer_idx
991
+ if config.is_kda_layer(layer_idx):
992
+ self.is_linear_attn = True
993
+ self.self_attn = KimiDeltaAttention(config=config, layer_idx=layer_idx)
994
+ elif config.is_mla:
995
+ self.is_linear_attn = False
996
+ self.self_attn = KimiMLAAttention(config=config, layer_idx=layer_idx)
997
+ else:
998
+ raise NotImplementedError
999
+ if (
1000
+ config.num_experts is not None
1001
+ and layer_idx >= config.first_k_dense_replace
1002
+ and layer_idx % getattr(config, "moe_layer_freq", 1) == 0
1003
+ ):
1004
+ self.block_sparse_moe = KimiSparseMoeBlock(config)
1005
+ else:
1006
+ self.mlp = KimiMLP(config)
1007
+ self.input_layernorm = KimiRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
1008
+ self.post_attention_layernorm = KimiRMSNorm(
1009
+ config.hidden_size, eps=config.rms_norm_eps
1010
+ )
1011
+
1012
+ # Attention residual
1013
+ self.use_attn_residuals = (
1014
+ getattr(config, "attn_res_block_size", None) is not None
1015
+ )
1016
+ if self.use_attn_residuals:
1017
+ self.attn_res_block_size = config.attn_res_block_size
1018
+ self.self_attention_res_norm = KimiRMSNorm(
1019
+ config.hidden_size, eps=config.rms_norm_eps
1020
+ )
1021
+ self.mlp_res_norm = KimiRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
1022
+ self.self_attention_res_proj = nn.Linear(config.hidden_size, 1, bias=False)
1023
+ self.mlp_res_proj = nn.Linear(config.hidden_size, 1, bias=False)
1024
+
1025
+ def forward(
1026
+ self,
1027
+ hidden_states: torch.Tensor,
1028
+ attention_mask: torch.Tensor | None = None,
1029
+ position_ids: torch.LongTensor | None = None,
1030
+ past_key_values: tuple[torch.Tensor] | None = None,
1031
+ output_attentions: bool | None = False,
1032
+ use_cache: bool | None = False,
1033
+ block_residual: torch.Tensor | None = None,
1034
+ **kwargs: Unpack[FlashAttentionKwargs],
1035
+ ):
1036
+ if self.use_attn_residuals:
1037
+ return self._forward_attn_residual(
1038
+ hidden_states,
1039
+ attention_mask,
1040
+ position_ids,
1041
+ past_key_values,
1042
+ output_attentions,
1043
+ use_cache,
1044
+ block_residual,
1045
+ **kwargs,
1046
+ )
1047
+
1048
+ residual = hidden_states
1049
+
1050
+ hidden_states = self.input_layernorm(hidden_states)
1051
+
1052
+ # Self Attention
1053
+ if self.is_linear_attn is False:
1054
+ hidden_states = self.self_attn(
1055
+ hidden_states=hidden_states,
1056
+ attention_mask=attention_mask,
1057
+ position_ids=position_ids,
1058
+ past_key_values=past_key_values,
1059
+ output_attentions=output_attentions,
1060
+ use_cache=use_cache,
1061
+ **kwargs,
1062
+ )
1063
+ else:
1064
+ hidden_states = self.self_attn(
1065
+ hidden_states=hidden_states,
1066
+ attention_mask=attention_mask,
1067
+ cache_params=past_key_values,
1068
+ output_attentions=output_attentions,
1069
+ use_cache=use_cache,
1070
+ **kwargs,
1071
+ )
1072
+ hidden_states = residual + hidden_states
1073
+
1074
+ # Fully Connected
1075
+ residual = hidden_states
1076
+ hidden_states = self.post_attention_layernorm(hidden_states)
1077
+ if hasattr(self, "block_sparse_moe"):
1078
+ hidden_states = self.block_sparse_moe(hidden_states)
1079
+ else:
1080
+ hidden_states = self.mlp(hidden_states)
1081
+ hidden_states = residual + hidden_states
1082
+
1083
+ return hidden_states
1084
+
1085
+ def _forward_attn_residual(
1086
+ self,
1087
+ hidden_states: torch.Tensor,
1088
+ attention_mask: torch.Tensor | None = None,
1089
+ position_ids: torch.LongTensor | None = None,
1090
+ past_key_values: tuple[torch.Tensor] | None = None,
1091
+ output_attentions: bool | None = False,
1092
+ use_cache: bool | None = False,
1093
+ block_residual: torch.Tensor | None = None,
1094
+ **kwargs: Unpack[FlashAttentionKwargs],
1095
+ ):
1096
+ batch_size, seq_len, hidden_size = hidden_states.shape
1097
+ prefix_sum = hidden_states
1098
+
1099
+ if block_residual is not None and block_residual.shape[1] > 0:
1100
+ hidden_states = _apply_attn_res(
1101
+ prefix_sum.view(-1, hidden_size),
1102
+ block_residual,
1103
+ self.self_attention_res_proj,
1104
+ self.self_attention_res_norm,
1105
+ ).view(batch_size, seq_len, hidden_size)
1106
+
1107
+ if self.layer_idx % self.attn_res_block_size == 0:
1108
+ block_residual = torch.cat(
1109
+ [block_residual, prefix_sum.view(-1, hidden_size).unsqueeze(1)], dim=1
1110
+ )
1111
+ prefix_sum = None
1112
+
1113
+ hidden_states = self.input_layernorm(hidden_states)
1114
+
1115
+ # Self Attention
1116
+ if self.is_linear_attn is False:
1117
+ hidden_states = self.self_attn(
1118
+ hidden_states=hidden_states,
1119
+ attention_mask=attention_mask,
1120
+ position_ids=position_ids,
1121
+ past_key_values=past_key_values,
1122
+ output_attentions=output_attentions,
1123
+ use_cache=use_cache,
1124
+ **kwargs,
1125
+ )
1126
+ else:
1127
+ hidden_states = self.self_attn(
1128
+ hidden_states=hidden_states,
1129
+ attention_mask=attention_mask,
1130
+ cache_params=past_key_values,
1131
+ output_attentions=output_attentions,
1132
+ use_cache=use_cache,
1133
+ **kwargs,
1134
+ )
1135
+
1136
+ if prefix_sum is not None:
1137
+ prefix_sum = prefix_sum + hidden_states
1138
+ else:
1139
+ prefix_sum = hidden_states
1140
+
1141
+ hidden_states = _apply_attn_res(
1142
+ prefix_sum.view(-1, hidden_size),
1143
+ block_residual,
1144
+ self.mlp_res_proj,
1145
+ self.mlp_res_norm,
1146
+ ).view(batch_size, seq_len, hidden_size)
1147
+
1148
+ hidden_states = self.post_attention_layernorm(hidden_states)
1149
+ if hasattr(self, "block_sparse_moe"):
1150
+ hidden_states = self.block_sparse_moe(hidden_states)
1151
+ else:
1152
+ hidden_states = self.mlp(hidden_states)
1153
+
1154
+ if prefix_sum is None:
1155
+ prefix_sum = hidden_states
1156
+ else:
1157
+ prefix_sum = prefix_sum + hidden_states
1158
+
1159
+ return prefix_sum, block_residual
1160
+
1161
+
1162
+ class KimiPreTrainedModel(PreTrainedModel):
1163
+ config_class = KimiLinearConfig
1164
+ base_model_prefix = "model"
1165
+ supports_gradient_checkpointing = True
1166
+ _no_split_modules = ["KimiDecoderLayer"]
1167
+ _skip_keys_device_placement = "past_key_values"
1168
+ _supports_flash_attn_2 = True
1169
+ _can_record_outputs = {
1170
+ "router_logits": OutputRecorder(KimiBlockSparseMLP, index=1),
1171
+ "hidden_states": KimiDecoderLayer,
1172
+ "attentions": KimiMLAAttention,
1173
+ }
1174
+ _is_stateful = True
1175
+
1176
+ def _init_weights(self, module):
1177
+ # HOTFIX: disk offloading attempts to initialize the meta tensors
1178
+ # but this is bad programming: we shouldn't be initializing these
1179
+ # params in the first place
1180
+ # the init attempt attempts to get `module.weight`, which DNE for qmodels
1181
+
1182
+ return
1183
+
1184
+ std = self.config.initializer_range
1185
+ if isinstance(module, nn.Linear):
1186
+ module.weight.data.normal_(mean=0.0, std=std)
1187
+ if module.bias is not None:
1188
+ module.bias.data.zero_()
1189
+ elif isinstance(module, nn.Embedding):
1190
+ module.weight.data.normal_(mean=0.0, std=std)
1191
+ if module.padding_idx is not None:
1192
+ module.weight.data[module.padding_idx].zero_()
1193
+
1194
+
1195
+ def _apply_attn_res(prefix_sum, block_residual, proj, norm):
1196
+ """
1197
+ prefix_sum: (num_tokens, hidden_size)
1198
+ block_residual: (num_tokens, num_blocks, hidden_size)
1199
+ """
1200
+ v = torch.cat((block_residual, prefix_sum.unsqueeze(1)), dim=1)
1201
+ v_float = v.float()
1202
+ variance = v_float.pow(2).mean(-1, keepdim=True)
1203
+ k = v_float * torch.rsqrt(variance + norm.variance_epsilon)
1204
+ score_weight = norm.weight.float() * proj.weight.squeeze(0).float()
1205
+ scores = (k * score_weight).sum(-1)
1206
+ probs = scores.softmax(-1).unsqueeze(1)
1207
+ hidden_states = torch.matmul(probs, v_float).squeeze(1)
1208
+ return hidden_states.to(v.dtype)
1209
+
1210
+
1211
+ class KimiLinearModel(KimiPreTrainedModel):
1212
+ def __init__(self, config: KimiLinearConfig):
1213
+ super().__init__(config)
1214
+ self.padding_idx = config.pad_token_id
1215
+ self.vocab_size = config.vocab_size
1216
+
1217
+ self.embed_tokens = nn.Embedding(
1218
+ config.vocab_size, config.hidden_size, self.padding_idx
1219
+ )
1220
+ self.layers = nn.ModuleList(
1221
+ [
1222
+ KimiDecoderLayer(config, layer_idx)
1223
+ for layer_idx in range(config.num_hidden_layers)
1224
+ ]
1225
+ )
1226
+ self.norm = KimiRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
1227
+
1228
+ self.use_attn_residuals = (
1229
+ getattr(config, "attn_res_block_size", None) is not None
1230
+ )
1231
+ if self.use_attn_residuals:
1232
+ self.output_attn_res_norm = KimiRMSNorm(
1233
+ config.hidden_size, eps=config.rms_norm_eps
1234
+ )
1235
+ self.output_attn_res_proj = nn.Linear(config.hidden_size, 1, bias=False)
1236
+
1237
+ from transformers.utils import is_flash_attn_2_available as _fa2_avail
1238
+
1239
+ _requested = getattr(config, "_attn_implementation", None)
1240
+ if _requested not in (None, "flash_attention_2") or not _fa2_avail():
1241
+ # Fall back gracefully when flash-attn2 is unavailable or a different impl is requested
1242
+ if _requested == "flash_attention_2" and not _fa2_avail():
1243
+ logger.warning_once(
1244
+ "flash_attention_2 requested but not available; falling back to sdpa."
1245
+ )
1246
+ config._attn_implementation = (
1247
+ _requested if _requested not in (None, "flash_attention_2") else "eager"
1248
+ )
1249
+ else:
1250
+ config._attn_implementation = "flash_attention_2"
1251
+
1252
+ self._use_flash_attention_2 = config._attn_implementation == "flash_attention_2"
1253
+ self.gradient_checkpointing = False
1254
+ # Initialize weights and apply final processing
1255
+ self.post_init()
1256
+
1257
+ def _update_linear_attn_mask(self, attention_mask, cache_position):
1258
+ """
1259
+ NOTE: Left-padding is used for linear attention mask.
1260
+ No need for zeroing states when
1261
+ 1. Cached forward
1262
+ 2. Attending to all inputs
1263
+ """
1264
+ linear_attn_mask = attention_mask
1265
+ if cache_position[0] > 0 or (
1266
+ attention_mask is not None and torch.all(attention_mask == 1)
1267
+ ):
1268
+ linear_attn_mask = None
1269
+ return linear_attn_mask
1270
+
1271
+ @check_model_inputs
1272
+ # @auto_docstring
1273
+ def forward(
1274
+ self,
1275
+ input_ids: torch.LongTensor = None,
1276
+ attention_mask: torch.Tensor | None = None,
1277
+ position_ids: torch.LongTensor | None = None,
1278
+ past_key_values: Cache | None = None,
1279
+ inputs_embeds: torch.FloatTensor | None = None,
1280
+ cache_position: torch.LongTensor | None = None,
1281
+ use_cache: bool | None = None,
1282
+ **kwargs: Unpack[TransformersKwargs],
1283
+ ) -> tuple | BaseModelOutputWithPast:
1284
+ use_cache = use_cache if use_cache is not None else self.config.use_cache
1285
+
1286
+ if (input_ids is None) and (inputs_embeds is None):
1287
+ raise ValueError(
1288
+ "You must specify exactly one of input_ids or inputs_embeds"
1289
+ )
1290
+
1291
+ # Get inputs_embeds
1292
+ if inputs_embeds is None:
1293
+ inputs_embeds = self.embed_tokens(input_ids)
1294
+
1295
+ if use_cache and past_key_values is None:
1296
+ past_key_values = KimiDynamicCache(config=self.config)
1297
+
1298
+ if cache_position is None:
1299
+ past_seen_tokens = (
1300
+ past_key_values.get_seq_length() if past_key_values is not None else 0
1301
+ )
1302
+ cache_position: torch.Tensor = torch.arange(
1303
+ past_seen_tokens,
1304
+ past_seen_tokens + inputs_embeds.shape[1],
1305
+ device=inputs_embeds.device,
1306
+ )
1307
+
1308
+ if position_ids is None:
1309
+ position_ids = cache_position.unsqueeze(0)
1310
+
1311
+ causal_mask = create_causal_mask(
1312
+ config=self.config,
1313
+ inputs_embeds=inputs_embeds,
1314
+ attention_mask=attention_mask,
1315
+ past_key_values=past_key_values,
1316
+ position_ids=position_ids,
1317
+ )
1318
+ linear_attn_mask = self._update_linear_attn_mask(attention_mask, cache_position)
1319
+
1320
+ hidden_states = inputs_embeds
1321
+ if past_key_values is not None:
1322
+ assert isinstance(past_key_values, KimiDynamicCache)
1323
+
1324
+ block_residual = None
1325
+ if self.use_attn_residuals:
1326
+ block_residual = hidden_states.new_zeros(
1327
+ hidden_states.shape[0] * hidden_states.shape[1],
1328
+ 0,
1329
+ hidden_states.shape[2],
1330
+ )
1331
+
1332
+ for decoder_layer in self.layers:
1333
+ layer_mask = (
1334
+ linear_attn_mask if decoder_layer.is_linear_attn else causal_mask
1335
+ )
1336
+
1337
+ if self.use_attn_residuals:
1338
+ hidden_states, block_residual = decoder_layer(
1339
+ hidden_states,
1340
+ attention_mask=layer_mask,
1341
+ past_key_values=past_key_values,
1342
+ cache_position=cache_position,
1343
+ block_residual=block_residual,
1344
+ **kwargs,
1345
+ )
1346
+ else:
1347
+ hidden_states = decoder_layer(
1348
+ hidden_states,
1349
+ attention_mask=layer_mask,
1350
+ past_key_values=past_key_values,
1351
+ cache_position=cache_position,
1352
+ **kwargs,
1353
+ )
1354
+
1355
+ if self.use_attn_residuals:
1356
+ hidden_states = self._apply_output_attn_res(hidden_states, block_residual)
1357
+
1358
+ hidden_states = self.norm(hidden_states)
1359
+
1360
+ return BaseModelOutputWithPast(
1361
+ last_hidden_state=hidden_states,
1362
+ past_key_values=past_key_values,
1363
+ )
1364
+
1365
+ def _apply_output_attn_res(self, hidden_states, block_residual):
1366
+ batch_size, seq_len, hidden_size = hidden_states.shape
1367
+ return _apply_attn_res(
1368
+ hidden_states.view(-1, hidden_size),
1369
+ block_residual,
1370
+ self.output_attn_res_proj,
1371
+ self.output_attn_res_norm,
1372
+ ).view(batch_size, seq_len, hidden_size)
1373
+
1374
+
1375
+ class KimiLinearForCausalLM(KimiPreTrainedModel, GenerationMixin):
1376
+ @classmethod
1377
+ def _supports_default_dynamic_cache(cls) -> bool:
1378
+ return False
1379
+
1380
+ _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"}
1381
+
1382
+ def __init__(self, config):
1383
+ super().__init__(config)
1384
+ self.model = KimiLinearModel(config)
1385
+ self.vocab_size = config.vocab_size
1386
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
1387
+
1388
+ # Initialize weights and apply final processing
1389
+ self.post_init()
1390
+
1391
+ @can_return_tuple
1392
+ # @auto_docstring
1393
+ def forward(
1394
+ self,
1395
+ input_ids: torch.LongTensor = None,
1396
+ attention_mask: torch.Tensor | None = None,
1397
+ position_ids: torch.LongTensor | None = None,
1398
+ past_key_values: list[torch.FloatTensor] | None = None,
1399
+ inputs_embeds: torch.FloatTensor | None = None,
1400
+ labels: torch.LongTensor | None = None,
1401
+ use_cache: bool | None = None,
1402
+ output_attentions: bool | None = None,
1403
+ output_hidden_states: bool | None = None,
1404
+ generation_mode: bool | None = None,
1405
+ return_dict: bool | None = None,
1406
+ cache_position: torch.LongTensor | None = None,
1407
+ **kwargs: Unpack[TransformersKwargs],
1408
+ ) -> tuple | CausalLMOutputWithPast:
1409
+ r"""
1410
+ Args:
1411
+ labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
1412
+ Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
1413
+ config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
1414
+ (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
1415
+ """
1416
+
1417
+ output_attentions = (
1418
+ output_attentions
1419
+ if output_attentions is not None
1420
+ else self.config.output_attentions
1421
+ )
1422
+ output_hidden_states = (
1423
+ output_hidden_states
1424
+ if output_hidden_states is not None
1425
+ else self.config.output_hidden_states
1426
+ )
1427
+ return_dict = (
1428
+ return_dict if return_dict is not None else self.config.use_return_dict
1429
+ )
1430
+
1431
+ outputs = self.model(
1432
+ input_ids=input_ids,
1433
+ attention_mask=attention_mask,
1434
+ position_ids=position_ids,
1435
+ past_key_values=past_key_values,
1436
+ inputs_embeds=inputs_embeds,
1437
+ use_cache=use_cache,
1438
+ output_attentions=output_attentions,
1439
+ output_hidden_states=output_hidden_states,
1440
+ return_dict=return_dict,
1441
+ cache_position=cache_position,
1442
+ )
1443
+
1444
+ logits = outputs[0]
1445
+ if generation_mode:
1446
+ logits = logits[:, -1:]
1447
+ logits = self.lm_head(logits)
1448
+
1449
+ loss = None
1450
+ if labels is not None:
1451
+ loss = self.loss_function(logits, labels, self.vocab_size, **kwargs)
1452
+
1453
+ return CausalLMOutputWithPast(
1454
+ loss=loss,
1455
+ logits=logits,
1456
+ past_key_values=outputs.past_key_values,
1457
+ hidden_states=outputs.hidden_states,
1458
+ attentions=outputs.attentions,
1459
+ )
modeling_kimi_linear.py ADDED
@@ -0,0 +1 @@
 
 
1
+ from .modeling_kimi_k3_linear import *
recipe.yaml ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ default_stage:
2
+ default_modifiers:
3
+ QuantizationModifier:
4
+ targets: [Linear]
5
+ ignore: ['re:.*self_attn.*', 're:.*shared_experts.*', 're:.*lm_head.*', 're:.*vision_tower.*']
6
+ scheme: MXFP4A16
7
+ bypass_divisibility_checks: false
tiktoken.model ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:b6c497a7469b33ced9c38afb1ad6e47f03f5e5dc05f15930799210ec050c5103
3
+ size 2795286
tokenization_kimi.py ADDED
@@ -0,0 +1,430 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ from logging import getLogger
3
+ from pathlib import Path
4
+ from shutil import copyfile
5
+ from typing import Dict, Iterator, List, Optional, Tuple, Union, cast
6
+
7
+ import tiktoken
8
+ from tiktoken.load import load_tiktoken_bpe
9
+ from tokenizers import AddedToken
10
+ from transformers.convert_slow_tokenizer import bytes_to_unicode
11
+ from transformers.tokenization_utils import PreTrainedTokenizer
12
+
13
+ try:
14
+ from .encoding_k3 import build_chat_segments, is_batched_conversation
15
+ except ImportError: # pragma: no cover - supports direct file execution/import.
16
+ from encoding_k3 import build_chat_segments, is_batched_conversation
17
+
18
+ logger = getLogger(__name__)
19
+ VOCAB_FILES_NAMES = {"vocab_file": "tiktoken.model"}
20
+
21
+
22
+ class TikTokenTokenizer(PreTrainedTokenizer):
23
+ """
24
+ Tokenizing and encoding/decoding text using the Tiktoken tokenizer. See megatron/tokenizer/tiktoken_tokenizer.py.
25
+
26
+ This tokenizer inherits from [`PreTrainedTokenizer`] which contains most of the main methods. Users should refer to
27
+ this superclass for more information regarding those methods.
28
+
29
+ Args:
30
+ vocab_file (`str`):
31
+ The path to the Tiktoken model file.
32
+ bos_token (`str` or `tokenizers.AddedToken`, *optional*, defaults to `"<|begin_of_text|>",`):
33
+ The beginning of sequence token that was used during pretraining. Can be used a sequence classifier token.
34
+ eos_token (`str` or `tokenizers.AddedToken`, *optional*, defaults to `"<|end_of_text|>"`):
35
+ The end of sequence token.
36
+ unk_token (`str` or `tokenizers.AddedToken`, *optional*, defaults to `"<|reserved_special_token_249|>"`):
37
+ The unknown token. A token that is not in the vocabulary cannot be converted to an ID and is set to be this
38
+ token instead. The second to last item in special_tokens.
39
+ pad_token (`str` or `tokenizers.AddedToken`, *optional*, defaults to `"<|reserved_special_token_250|>"`):
40
+ The token used for padding, for example when batching sequences of different lengths.
41
+ additional_special_tokens (list of `str`, *optional*):
42
+ A tuple or a list of additional tokens, which will be marked as `special`, meaning that they will be
43
+ skipped when decoding if `skip_special_tokens` is set to `True`.
44
+ """
45
+
46
+ vocab_files_names = VOCAB_FILES_NAMES
47
+
48
+ model_input_names = ["input_ids", "attention_mask"]
49
+
50
+ special_tokens: Dict[str, int]
51
+
52
+ num_reserved_special_tokens = 256
53
+
54
+ pat_str = "|".join(
55
+ [
56
+ r"""[\p{Han}]+""",
57
+ r"""[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}&&[^\p{Han}]]*[\p{Ll}\p{Lm}\p{Lo}\p{M}&&[^\p{Han}]]+(?i:'s|'t|'re|'ve|'m|'ll|'d)?""",
58
+ r"""[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}&&[^\p{Han}]]+[\p{Ll}\p{Lm}\p{Lo}\p{M}&&[^\p{Han}]]*(?i:'s|'t|'re|'ve|'m|'ll|'d)?""",
59
+ r"""\p{N}{1,3}""",
60
+ r""" ?[^\s\p{L}\p{N}]+[\r\n]*""",
61
+ r"""\s*[\r\n]+""",
62
+ r"""\s+(?!\S)""",
63
+ r"""\s+""",
64
+ ]
65
+ )
66
+
67
+ def __init__(
68
+ self,
69
+ vocab_file,
70
+ bos_token: Union[str, AddedToken] = "[BOS]",
71
+ eos_token: Union[str, AddedToken] = "[EOS]",
72
+ unk_token: Union[str, AddedToken, None] = None,
73
+ pad_token: Union[str, AddedToken, None] = None,
74
+ additional_special_tokens: List[str] = None,
75
+ added_tokens_decoder: Optional[dict] = None,
76
+ **kwargs,
77
+ ):
78
+ assert os.path.isfile(vocab_file), vocab_file
79
+
80
+ if additional_special_tokens is None:
81
+ additional_special_tokens = [
82
+ "<|im_end|>",
83
+ "<|im_user|>",
84
+ "<|im_assistant|>",
85
+ "<|start_header_id|>",
86
+ "<|end_header_id|>",
87
+ "[EOT]",
88
+ "<|im_system|>",
89
+ "<|im_middle|>",
90
+ ]
91
+
92
+ if added_tokens_decoder:
93
+ special_tokens_mapping = {
94
+ i: added_tokens_decoder[i].content for i in added_tokens_decoder
95
+ }
96
+ else:
97
+ special_tokens_mapping = {}
98
+
99
+ self.vocab_file = vocab_file
100
+ mergeable_ranks = load_tiktoken_bpe(vocab_file)
101
+ num_base_tokens = len(mergeable_ranks)
102
+ self.special_tokens = {
103
+ special_tokens_mapping.get(i, f"<|reserved_token_{i}|>"): i
104
+ for i in range(
105
+ num_base_tokens, num_base_tokens + self.num_reserved_special_tokens
106
+ )
107
+ }
108
+
109
+ self.model = tiktoken.Encoding(
110
+ name=Path(vocab_file).name,
111
+ pat_str=self.pat_str,
112
+ mergeable_ranks=mergeable_ranks,
113
+ special_tokens=self.special_tokens,
114
+ )
115
+ logger.info(f"Reloaded tiktoken model from {vocab_file}")
116
+
117
+ self.n_words: int = self.model.n_vocab
118
+ # BOS / EOS token IDs
119
+ self.bos_id: int = self.special_tokens[str(bos_token)]
120
+ self.eos_id: int = self.special_tokens[str(eos_token)]
121
+ logger.info(
122
+ f"#words: {self.n_words} - BOS ID: {self.bos_id} - EOS ID: {self.eos_id}"
123
+ )
124
+
125
+ self.pad_id: int = self.special_tokens[str(pad_token)]
126
+ self.unk_id: int = self.special_tokens[str(unk_token)]
127
+
128
+ self.byte_encoder = bytes_to_unicode()
129
+ self.byte_decoder = {v: k for k, v in self.byte_encoder.items()}
130
+
131
+ self.decoder = {}
132
+ for i in range(self.n_words):
133
+ # Taken from https://gist.github.com/xenova/a452a6474428de0182b17605a98631ee
134
+ decoding = "".join(
135
+ [
136
+ self.byte_encoder[ord(char)]
137
+ for char in self.model.decode_single_token_bytes(i).decode(
138
+ "latin-1"
139
+ )
140
+ ]
141
+ )
142
+ self.decoder[i] = decoding
143
+
144
+ self.encoder = {}
145
+ for i in range(self.n_words):
146
+ if i in self.decoder:
147
+ self.encoder[self.decoder[i]] = i
148
+
149
+ super().__init__(
150
+ bos_token=bos_token,
151
+ eos_token=eos_token,
152
+ unk_token=unk_token,
153
+ pad_token=pad_token,
154
+ additional_special_tokens=additional_special_tokens,
155
+ added_tokens_decoder=added_tokens_decoder,
156
+ **kwargs,
157
+ )
158
+ self.all_special_ids_set = set(self.all_special_ids)
159
+
160
+ def _encode_text_piece(
161
+ self, text: str, allow_special_tokens: bool = True
162
+ ) -> List[int]:
163
+ # The tiktoken tokenizer can handle <=400k chars without
164
+ # pyo3_runtime.PanicException.
165
+ TIKTOKEN_MAX_ENCODE_CHARS = 400_000
166
+
167
+ # https://github.com/openai/tiktoken/issues/195
168
+ # Here we iterate over subsequences and split if we exceed the limit
169
+ # of max consecutive non-whitespace or whitespace characters.
170
+ MAX_NO_WHITESPACES_CHARS = 25_000
171
+
172
+ t: List[int] = []
173
+ for i in range(0, len(text), TIKTOKEN_MAX_ENCODE_CHARS):
174
+ for substr in self._split_whitespaces_or_nonwhitespaces(
175
+ text[i : i + TIKTOKEN_MAX_ENCODE_CHARS],
176
+ MAX_NO_WHITESPACES_CHARS,
177
+ ):
178
+ if allow_special_tokens:
179
+ t.extend(
180
+ # structural markers: encode <|...|> as their special token IDs
181
+ self.model.encode(
182
+ substr,
183
+ allowed_special="all",
184
+ )
185
+ )
186
+ else:
187
+ t.extend(
188
+ # user/tool text: encode any <|...|> as ordinary BPE tokens (never as control tokens)
189
+ self.model.encode(
190
+ substr,
191
+ disallowed_special=(),
192
+ )
193
+ )
194
+
195
+ return t
196
+
197
+ def encode(
198
+ self, text: str, allow_special_tokens: bool = True, **kwargs
199
+ ) -> List[int]:
200
+ """
201
+ Encodes a string into a list of token IDs.
202
+
203
+ Args:
204
+ text (str): The input string to be encoded.
205
+
206
+ Returns:
207
+ list[int]: A list of token IDs.
208
+ """
209
+ # If there are other args, we should call super().encode because there are a lot of code
210
+ # to handle those args. supper().encode finally will call _tokenize and _convert_token_to_id.
211
+ # NOTE: our encode method is not compatible with the super().encode method,
212
+ # e.g. split_special_tokens' default is True in our encode method.
213
+ if len(kwargs) > 0:
214
+ logger.warning(f"Calling super().encode with {kwargs}")
215
+ return super().encode(text, **kwargs)
216
+
217
+ assert type(text) is str
218
+ return self._encode_text_piece(text, allow_special_tokens=allow_special_tokens)
219
+
220
+ def decode(self, token_ids: Union[int, List[int]], **kwargs) -> str:
221
+ """
222
+ Decodes a list of token IDs into a string.
223
+
224
+ Args:
225
+ token_ids (List[int]): The list of token IDs to be decoded.
226
+
227
+ Returns:
228
+ str: The decoded string.
229
+ """
230
+ # If there are other args, we should call super().decode because there are a lot of code
231
+ # to handle those args. supper().encode finally will call convert_tokens_to_string and _convert_id_to_token.
232
+ if len(kwargs) > 0:
233
+ return super().decode(token_ids, **kwargs)
234
+
235
+ if type(token_ids) is int:
236
+ token_ids = [token_ids]
237
+
238
+ return self.model.decode(cast(List[int], token_ids))
239
+
240
+ @staticmethod
241
+ def _split_whitespaces_or_nonwhitespaces(
242
+ s: str, max_consecutive_slice_len: int
243
+ ) -> Iterator[str]:
244
+ """
245
+ Splits the string `s` so that each substring contains no more than `max_consecutive_slice_len`
246
+ consecutive whitespaces or consecutive non-whitespaces.
247
+ """
248
+ current_slice_len = 0
249
+ current_slice_is_space = s[0].isspace() if len(s) > 0 else False
250
+ slice_start = 0
251
+
252
+ for i in range(len(s)):
253
+ is_now_space = s[i].isspace()
254
+
255
+ if current_slice_is_space ^ is_now_space:
256
+ current_slice_len = 1
257
+ current_slice_is_space = is_now_space
258
+ else:
259
+ current_slice_len += 1
260
+ if current_slice_len > max_consecutive_slice_len:
261
+ yield s[slice_start:i]
262
+ slice_start = i
263
+ current_slice_len = 1
264
+ yield s[slice_start:]
265
+
266
+ def _encode_chat_segments(self, segments) -> List[int]:
267
+ token_ids: List[int] = []
268
+ for segment in segments:
269
+ token_ids.extend(
270
+ self._encode_text_piece(
271
+ segment.text,
272
+ allow_special_tokens=segment.allow_special,
273
+ )
274
+ )
275
+ return token_ids
276
+
277
+ @staticmethod
278
+ def _truncate(
279
+ ids: List[int], truncation: bool = False, max_length: Optional[int] = None
280
+ ) -> List[int]:
281
+ if truncation and max_length is not None:
282
+ return ids[:max_length]
283
+ return ids
284
+
285
+ def _format_chat_token_output(
286
+ self,
287
+ encoded_inputs: List[List[int]],
288
+ *,
289
+ is_batched: bool,
290
+ padding=False,
291
+ truncation: bool = False,
292
+ max_length: Optional[int] = None,
293
+ return_tensors=None,
294
+ return_dict: bool = False,
295
+ ):
296
+ encoded_inputs = [
297
+ self._truncate(ids, truncation=truncation, max_length=max_length)
298
+ for ids in encoded_inputs
299
+ ]
300
+
301
+ needs_batch_encoding = (
302
+ is_batched or padding or return_tensors is not None or return_dict
303
+ )
304
+ if not needs_batch_encoding:
305
+ return encoded_inputs[0]
306
+
307
+ features = [
308
+ {"input_ids": ids, "attention_mask": [1] * len(ids)}
309
+ for ids in encoded_inputs
310
+ ]
311
+ batch = self.pad(
312
+ features,
313
+ padding=padding,
314
+ max_length=max_length if padding else None,
315
+ return_attention_mask=True,
316
+ return_tensors=return_tensors,
317
+ )
318
+
319
+ if return_dict:
320
+ return batch
321
+ if is_batched:
322
+ return batch["input_ids"]
323
+ return batch["input_ids"][0] if return_tensors is None else batch["input_ids"]
324
+
325
+ """ ----- Below are the abstract methods required by PreTrainedTokenizer ----- """
326
+
327
+ @property
328
+ def vocab_size(self) -> int:
329
+ return self.n_words
330
+
331
+ def get_vocab(self) -> Dict[str, int]:
332
+ return self.encoder
333
+
334
+ def _tokenize(self, text: str, **kwargs) -> List[str]:
335
+ return [self.decoder[t] for t in self.encode(text)]
336
+
337
+ def _convert_token_to_id(self, token: str) -> int:
338
+ return self.encoder.get(token, self.unk_id)
339
+
340
+ def _convert_id_to_token(self, index: int) -> str:
341
+ return self.decoder.get(index)
342
+
343
+ @staticmethod
344
+ def clean_up_tokenization(out_string: str) -> str:
345
+ return out_string
346
+
347
+ def convert_tokens_to_string(self, tokens: List[str]) -> str:
348
+ text = "".join(tokens)
349
+ text = bytearray([self.byte_decoder[c] for c in text]).decode(
350
+ "utf-8", "replace"
351
+ )
352
+ return text
353
+
354
+ def save_vocabulary(
355
+ self, save_directory: str, filename_prefix: Optional[str] = None
356
+ ) -> Tuple[str]:
357
+ if not os.path.isdir(save_directory):
358
+ raise ValueError(
359
+ f"vocabulary path ({save_directory}) should be a directory"
360
+ )
361
+ out_vocab_file = os.path.join(
362
+ save_directory,
363
+ (filename_prefix + "-" if filename_prefix else "")
364
+ + VOCAB_FILES_NAMES["vocab_file"],
365
+ )
366
+
367
+ if os.path.abspath(self.vocab_file) != os.path.abspath(
368
+ out_vocab_file
369
+ ) and os.path.isfile(self.vocab_file):
370
+ copyfile(self.vocab_file, out_vocab_file)
371
+
372
+ return (out_vocab_file,)
373
+
374
+ def apply_chat_template(
375
+ self,
376
+ conversation,
377
+ tools: Optional[list[dict]] = None,
378
+ tokenize: bool = False,
379
+ add_generation_prompt: bool = True,
380
+ thinking: bool = True,
381
+ padding=False,
382
+ truncation: bool = False,
383
+ max_length: Optional[int] = None,
384
+ return_tensors=None,
385
+ return_dict: bool = False,
386
+ **kwargs,
387
+ ):
388
+ # Tokenizer-level rendering reorders tool result messages to match
389
+ # assistant tool_calls, normalizes per-call arguments and response
390
+ # schema, then encodes the resulting XTML structure segment-by-segment.
391
+ is_batched = is_batched_conversation(conversation)
392
+ conversations = conversation if is_batched else [conversation]
393
+ image_prompts = kwargs.pop("image_prompts", None)
394
+ if is_batched and image_prompts is not None:
395
+ raise ValueError("image_prompts is only supported for one chat.")
396
+
397
+ # by default set thinking effort to max
398
+ kwargs.setdefault("thinking_effort", "max")
399
+
400
+ segment_batches = [
401
+ build_chat_segments(
402
+ messages,
403
+ tools=tools,
404
+ add_generation_prompt=add_generation_prompt,
405
+ thinking=thinking,
406
+ image_prompts=image_prompts,
407
+ **kwargs,
408
+ )
409
+ for messages in conversations
410
+ ]
411
+
412
+ if not tokenize:
413
+ rendered = [
414
+ "".join(segment.text for segment in segments)
415
+ for segments in segment_batches
416
+ ]
417
+ return rendered if is_batched else rendered[0]
418
+
419
+ encoded_inputs = [
420
+ self._encode_chat_segments(segments) for segments in segment_batches
421
+ ]
422
+ return self._format_chat_token_output(
423
+ encoded_inputs,
424
+ is_batched=is_batched,
425
+ padding=padding,
426
+ truncation=truncation,
427
+ max_length=max_length,
428
+ return_tensors=return_tensors,
429
+ return_dict=return_dict,
430
+ )
tokenizer_config.json ADDED
@@ -0,0 +1,224 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "added_tokens_decoder": {
3
+ "163584": {
4
+ "content": "[BOS]",
5
+ "lstrip": false,
6
+ "normalized": false,
7
+ "rstrip": false,
8
+ "single_word": false,
9
+ "special": true
10
+ },
11
+ "163585": {
12
+ "content": "[EOS]",
13
+ "lstrip": false,
14
+ "normalized": false,
15
+ "rstrip": false,
16
+ "single_word": false,
17
+ "special": true
18
+ },
19
+ "163586": {
20
+ "content": "<|end_of_msg|>",
21
+ "lstrip": false,
22
+ "normalized": false,
23
+ "rstrip": false,
24
+ "single_word": false,
25
+ "special": true
26
+ },
27
+ "163587": {
28
+ "content": "<|open|>",
29
+ "lstrip": false,
30
+ "normalized": false,
31
+ "rstrip": false,
32
+ "single_word": false,
33
+ "special": false
34
+ },
35
+ "163588": {
36
+ "content": "<|close|>",
37
+ "lstrip": false,
38
+ "normalized": false,
39
+ "rstrip": false,
40
+ "single_word": false,
41
+ "special": false
42
+ },
43
+ "163589": {
44
+ "content": "<|sep|>",
45
+ "lstrip": false,
46
+ "normalized": false,
47
+ "rstrip": false,
48
+ "single_word": false,
49
+ "special": false
50
+ },
51
+ "163590": {
52
+ "content": "[start_header_id]",
53
+ "lstrip": false,
54
+ "normalized": false,
55
+ "rstrip": false,
56
+ "single_word": false,
57
+ "special": true
58
+ },
59
+ "163591": {
60
+ "content": "[end_header_id]",
61
+ "lstrip": false,
62
+ "normalized": false,
63
+ "rstrip": false,
64
+ "single_word": false,
65
+ "special": true
66
+ },
67
+ "163593": {
68
+ "content": "[EOT]",
69
+ "lstrip": false,
70
+ "normalized": false,
71
+ "rstrip": false,
72
+ "single_word": false,
73
+ "special": true
74
+ },
75
+ "163602": {
76
+ "content": "<|media_begin|>",
77
+ "lstrip": false,
78
+ "normalized": false,
79
+ "rstrip": false,
80
+ "single_word": false,
81
+ "special": true
82
+ },
83
+ "163603": {
84
+ "content": "<|media_content|>",
85
+ "lstrip": false,
86
+ "normalized": false,
87
+ "rstrip": false,
88
+ "single_word": false,
89
+ "special": true
90
+ },
91
+ "163604": {
92
+ "content": "<|media_end|>",
93
+ "lstrip": false,
94
+ "normalized": false,
95
+ "rstrip": false,
96
+ "single_word": false,
97
+ "special": true
98
+ },
99
+ "163605": {
100
+ "content": "<|media_pad|>",
101
+ "lstrip": false,
102
+ "normalized": false,
103
+ "rstrip": false,
104
+ "single_word": false,
105
+ "special": true
106
+ },
107
+ "163649": {
108
+ "content": "<osagent_mode>",
109
+ "lstrip": false,
110
+ "normalized": false,
111
+ "rstrip": false,
112
+ "single_word": false,
113
+ "special": true
114
+ },
115
+ "163838": {
116
+ "content": "[UNK]",
117
+ "lstrip": false,
118
+ "normalized": false,
119
+ "rstrip": false,
120
+ "single_word": false,
121
+ "special": true
122
+ },
123
+ "163839": {
124
+ "content": "[PAD]",
125
+ "lstrip": false,
126
+ "normalized": false,
127
+ "rstrip": false,
128
+ "single_word": false,
129
+ "special": true
130
+ },
131
+ "163840": {
132
+ "content": "<|im_end|>",
133
+ "lstrip": false,
134
+ "normalized": false,
135
+ "rstrip": false,
136
+ "single_word": false,
137
+ "special": true
138
+ },
139
+ "163841": {
140
+ "content": "<|im_user|>",
141
+ "lstrip": false,
142
+ "normalized": false,
143
+ "rstrip": false,
144
+ "single_word": false,
145
+ "special": true
146
+ },
147
+ "163842": {
148
+ "content": "<|im_assistant|>",
149
+ "lstrip": false,
150
+ "normalized": false,
151
+ "rstrip": false,
152
+ "single_word": false,
153
+ "special": true
154
+ },
155
+ "163843": {
156
+ "content": "<|start_header_id|>",
157
+ "lstrip": false,
158
+ "normalized": false,
159
+ "rstrip": false,
160
+ "single_word": false,
161
+ "special": true
162
+ },
163
+ "163844": {
164
+ "content": "<|end_header_id|>",
165
+ "lstrip": false,
166
+ "normalized": false,
167
+ "rstrip": false,
168
+ "single_word": false,
169
+ "special": true
170
+ },
171
+ "163845": {
172
+ "content": "<|im_system|>",
173
+ "lstrip": false,
174
+ "normalized": false,
175
+ "rstrip": false,
176
+ "single_word": false,
177
+ "special": true
178
+ },
179
+ "163846": {
180
+ "content": "<|im_middle|>",
181
+ "lstrip": false,
182
+ "normalized": false,
183
+ "rstrip": false,
184
+ "single_word": false,
185
+ "special": true
186
+ }
187
+ },
188
+ "additional_special_tokens": [
189
+ "<|im_end|>",
190
+ "<|im_user|>",
191
+ "<|im_assistant|>",
192
+ "<|start_header_id|>",
193
+ "<|end_header_id|>",
194
+ "[EOT]",
195
+ "<|im_system|>",
196
+ "<|im_middle|>"
197
+ ],
198
+ "auto_map": {
199
+ "AutoTokenizer": [
200
+ "tokenization_kimi.TikTokenTokenizer",
201
+ null
202
+ ]
203
+ },
204
+ "backend": "custom",
205
+ "bos_token": "[BOS]",
206
+ "clean_up_tokenization_spaces": false,
207
+ "eos_token": "[EOS]",
208
+ "extra_special_tokens": [
209
+ "<|im_end|>",
210
+ "<|im_user|>",
211
+ "<|im_assistant|>",
212
+ "<|start_header_id|>",
213
+ "<|end_header_id|>",
214
+ "[EOT]",
215
+ "<|im_system|>",
216
+ "<|im_middle|>"
217
+ ],
218
+ "is_local": false,
219
+ "local_files_only": false,
220
+ "model_max_length": 1000000000000000019884624838656,
221
+ "pad_token": "[PAD]",
222
+ "tokenizer_class": "TikTokenTokenizer",
223
+ "unk_token": "[UNK]"
224
+ }