Skip to content

Commit b2f13d7

Browse files
committed
feat(kimi-k3): Kimi K3 support: MXFP4 to BF16 conversion and inference graph
New arch kimi-k3 (KimiK3ForConditionalGeneration, text_config kimi_linear): - compressed-tensors mxfp4-pack-quantized dequant (E2M1 nibbles low=even, E8M0 uint8 scale, group 32), verified against real Kimi-K3 shard tensors - KDA safe gate: g = lower_bound * sigmoid(exp(A_log) * (f_b(f_a(x)) + dt_bias)), full-rank output gate g_proj, A_log sliced to [:n_head] at conversion - gated MLA (NoPE): attn * sigmoid(g_proj(x)) before o_proj, q_lora_rank 1536 - Stable LatentMoE: router on full hidden, experts in 3584 latent space, RMSNorm after the weighted expert sum, then up-projection - AttnRes: residual-stream snapshot bank every attn_res_block_size layers, softmax mixtures before attention, before FFN and at model output - SiTU-GLU activation (LLM_FFN_SITU): beta*tanh(g/beta)*sigmoid(g) * lbeta*tanh(up/lbeta) Tested on a synthetic mini K3: convert -> BF16 GGUF -> Q8_0 -> generation.
1 parent 0e4a036 commit b2f13d7

15 files changed

Lines changed: 910 additions & 3 deletions

conversion/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -121,6 +121,7 @@
121121
"JinaEmbeddingsV5Model": "bert",
122122
"KORMoForCausalLM": "qwen",
123123
"KimiK25ForConditionalGeneration": "deepseek",
124+
"KimiK3ForConditionalGeneration": "kimi_k3",
124125
"KimiLinearForCausalLM": "kimi_linear",
125126
"KimiLinearModel": "kimi_linear",
126127
"KimiVLForConditionalGeneration": "deepseek",

conversion/base.py

Lines changed: 41 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -386,6 +386,29 @@ def dequant_gptq(g_idx: Tensor, qweight: Tensor, qzeros: Tensor, scales: Tensor)
386386

387387
return (scales[g_idx].float() * (weight - zeros[g_idx]).float()).T
388388

389+
def dequant_mxfp4_packed(w: Tensor, scale: Tensor, group_size: int) -> Tensor:
390+
# compressed-tensors "mxfp4-pack-quantized":
391+
# w: uint8 [..., K/2], two FP4 (E2M1) values per byte, low nibble = even index
392+
# scale: uint8 [..., K/group_size], E8M0 exponent, scale = 2^(x - 127)
393+
assert w.dtype == torch.uint8
394+
assert scale.dtype == torch.uint8
395+
396+
kvalues = torch.tensor(
397+
[0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0],
398+
dtype=torch.float32,
399+
)
400+
if self.lazy:
401+
kvalues = LazyTorchTensor.from_eager(kvalues)
402+
403+
lo = (w & 0x0F).to(torch.long)
404+
hi = (w >> 4).to(torch.long)
405+
# interleave: [..., K/2, 2] -> [..., K]
406+
vals = torch.stack((kvalues[lo], kvalues[hi]), dim=-1).reshape(*w.shape[:-1], w.shape[-1] * 2)
407+
408+
exp = torch.ldexp(torch.ones_like(scale, dtype=torch.float32), scale.to(torch.int32) - 127)
409+
vals = vals.reshape(*vals.shape[:-1], -1, group_size) * exp.unsqueeze(-1)
410+
return vals.reshape(*w.shape[:-1], w.shape[-1] * 2)
411+
389412
def dequant_packed(w: Tensor, scale: Tensor, shape_tensor: Tensor, zero_point: Tensor | None, num_bits: int, group_size: int):
390413
assert w.dtype == torch.int32
391414
shape = tuple(shape_tensor.tolist())
@@ -531,6 +554,23 @@ def dequant_packed(w: Tensor, scale: Tensor, shape_tensor: Tensor, zero_point: T
531554
tensors_to_remove += [base_name + n for n in ("_packed", "_shape", "_scale")]
532555
if (base_name + "_zero_point") in self.model_tensors:
533556
tensors_to_remove.append(base_name + "_zero_point")
557+
elif quant_format == "mxfp4-pack-quantized":
558+
assert weight_config.get("strategy") == "group"
559+
assert weight_config.get("type") == "float"
560+
assert weight_config.get("num_bits") == 4
561+
group_size = weight_config.get("group_size")
562+
assert isinstance(group_size, int)
563+
for name in self.model_tensors.keys():
564+
if name.endswith(".weight_packed"):
565+
base_name = name.removesuffix("_packed")
566+
w = self.model_tensors[name]
567+
scale = self.model_tensors[base_name + "_scale"]
568+
new_tensors[base_name] = (
569+
lambda w=w, scale=scale: dequant_mxfp4_packed(w(), scale(), group_size)
570+
)
571+
tensors_to_remove += [base_name + n for n in ("_packed", "_scale")]
572+
if (base_name + "_shape") in self.model_tensors:
573+
tensors_to_remove.append(base_name + "_shape")
534574
elif nvfp4_compressed_tensors:
535575
# Don't error from compressed-tensors, we'll handle them in _generate_nvfp4_tensors
536576
pass
@@ -2633,7 +2673,7 @@ def get_model_architecture(hparams: dict[str, Any], model_type: ModelType) -> st
26332673
# Step3-VL keeps text config under text_config but uses a custom top-level architecture.
26342674
# For text conversion we route to a dedicated text-only class.
26352675
# TODO: refactor this later to avoid adding exception here
2636-
if model_type == ModelType.TEXT and arch in ("StepVLForConditionalGeneration", "Sarashina2VisionForCausalLM", "Exaone4_5_ForConditionalGeneration", "Step3p7ForConditionalGeneration"):
2676+
if model_type == ModelType.TEXT and arch in ("StepVLForConditionalGeneration", "Sarashina2VisionForCausalLM", "Exaone4_5_ForConditionalGeneration", "Step3p7ForConditionalGeneration", "KimiK3ForConditionalGeneration"):
26372677
return arch
26382678

26392679
# if "architectures" is found in the sub-config, use that instead

conversion/kimi_k3.py

Lines changed: 69 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,69 @@
1+
from __future__ import annotations
2+
3+
from typing import Iterable, TYPE_CHECKING
4+
5+
import torch
6+
7+
if TYPE_CHECKING:
8+
from torch import Tensor
9+
10+
from .base import ModelBase, gguf
11+
from .kimi_linear import KimiLinearModel
12+
13+
14+
@ModelBase.register("KimiK3ForConditionalGeneration")
15+
class KimiK3Model(KimiLinearModel):
16+
"""Kimi K3: hybrid KDA + gated-MLA (NoPE) with Attention Residuals and Stable LatentMoE.
17+
18+
Text config is `kimi_linear` with K3 extensions:
19+
- SiTU-GLU activation (soft-capped SiLU) in dense MLP, shared and routed experts
20+
- AttnRes: residual-stream snapshot bank every `attn_res_block_size` layers,
21+
softmax mixtures before attention, before MLP and at model output
22+
- Stable LatentMoE: routed experts run in a `routed_expert_hidden_size` latent
23+
space (down proj -> experts -> weighted sum -> RMSNorm -> up proj)
24+
- KDA safe gate: g_log = gate_lower_bound * sigmoid(exp(A_log) * (g_raw + dt_bias))
25+
with a full-rank output gate g_proj instead of the low-rank g_a/g_b pair
26+
- MLA output gate: attn = attn * sigmoid(g_proj(x)) before o_proj
27+
Routed expert weights are MXFP4 (compressed-tensors), dequantized in ModelBase.
28+
"""
29+
model_arch = gguf.MODEL_ARCH.KIMI_K3
30+
31+
def set_gguf_parameters(self):
32+
super().set_gguf_parameters()
33+
34+
# Stable LatentMoE
35+
self.gguf_writer.add_moe_latent_size(self.hparams["routed_expert_hidden_size"])
36+
37+
# SiTU-GLU activation parameters
38+
self.gguf_writer.add_situ_beta(self.hparams["activation_situ_beta"])
39+
self.gguf_writer.add_situ_linear_beta(self.hparams["activation_situ_linear_beta"])
40+
41+
# Attention residuals
42+
self.gguf_writer.add_attn_res_block_size(self.hparams["attn_res_block_size"])
43+
44+
# KDA safe gate lower bound
45+
self.gguf_writer.add_kda_gate_lower_bound(self.hparams["linear_attn_config"]["gate_lower_bound"])
46+
47+
def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:
48+
# text-only conversion: vision tensors are handled by the mmproj path
49+
if name.startswith(("vision_tower.", "mm_projector.")):
50+
return
51+
52+
name = name.removeprefix("language_model.")
53+
54+
# K3 checkpoints store A_log as [head_dim] (128) but only the first
55+
# num_heads (96) entries are used. The safe-gate formula is
56+
# g_log = gate_lower_bound * sigmoid(exp(A_log) * (g_raw + dt_bias))
57+
# so we store exp(A_log) directly (unlike Kimi-Linear's -exp(A_log)).
58+
if name.endswith(".A_log"):
59+
n_head = self.hparams["num_attention_heads"]
60+
data_torch = torch.exp(data_torch.float()[:n_head])
61+
# skip KimiLinearModel's -exp(A_log) handling
62+
yield from super(KimiLinearModel, self).modify_tensors(data_torch, name, bid)
63+
return
64+
65+
# res projections are stored as [1, n_embd]: flatten to [n_embd]
66+
if name.endswith(("_res_proj.weight", "_res_norm.weight")):
67+
data_torch = data_torch.reshape(-1)
68+
69+
yield from super().modify_tensors(data_torch, name, bid)

gguf-py/gguf/constants.py

Lines changed: 70 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -126,6 +126,9 @@ class LLM:
126126
EXPERTS_PER_GROUP = "{arch}.experts_per_group"
127127
MOE_EVERY_N_LAYERS = "{arch}.moe_every_n_layers"
128128
MOE_LATENT_SIZE = "{arch}.moe_latent_size"
129+
SITU_BETA = "{arch}.situ_beta"
130+
SITU_LINEAR_BETA = "{arch}.situ_linear_beta"
131+
ATTN_RES_BLOCK_SIZE = "{arch}.attn_res_block_size"
129132
NEXTN_PREDICT_LAYERS = "{arch}.nextn_predict_layers"
130133
NUM_DEEPSTACK_LAYERS = "{arch}.n_deepstack_layers"
131134
DEEPSTACK_MAPPING = "{arch}.deepstack_mapping"
@@ -243,7 +246,8 @@ class SSM:
243246
DT_B_C_RMS = "{arch}.ssm.dt_b_c_rms"
244247

245248
class KDA:
246-
HEAD_DIM = "{arch}.kda.head_dim"
249+
HEAD_DIM = "{arch}.kda.head_dim"
250+
GATE_LOWER_BOUND = "{arch}.kda.gate_lower_bound"
247251

248252
class WKV:
249253
HEAD_SIZE = "{arch}.wkv.head_size"
@@ -545,6 +549,7 @@ class MODEL_ARCH(IntEnum):
545549
LLAMA_EMBED = auto()
546550
MAINCODER = auto()
547551
KIMI_LINEAR = auto()
552+
KIMI_K3 = auto()
548553
TALKIE = auto()
549554
MELLUM = auto()
550555
NANBEIGE = auto()
@@ -620,6 +625,13 @@ class MODEL_TENSOR(IntEnum):
620625
FFN_GATE_TID2EID = auto()
621626
MOE_LATENT_DOWN = auto() # nemotron 3 super
622627
MOE_LATENT_UP = auto() # nemotron 3 super
628+
MOE_LATENT_NORM = auto() # kimi k3
629+
ATTN_RES_NORM = auto() # kimi k3
630+
ATTN_RES_PROJ = auto() # kimi k3
631+
FFN_RES_NORM = auto() # kimi k3
632+
FFN_RES_PROJ = auto() # kimi k3
633+
OUTPUT_RES_NORM = auto() # kimi k3
634+
OUTPUT_RES_PROJ = auto() # kimi k3
623635
ATTN_Q_NORM = auto()
624636
ATTN_K_NORM = auto()
625637
LAYER_OUT_NORM = auto()
@@ -1135,6 +1147,7 @@ class MODEL_TENSOR(IntEnum):
11351147
MODEL_ARCH.LLAMA_EMBED: "llama-embed",
11361148
MODEL_ARCH.MAINCODER: "maincoder",
11371149
MODEL_ARCH.KIMI_LINEAR: "kimi-linear",
1150+
MODEL_ARCH.KIMI_K3: "kimi-k3",
11381151
MODEL_ARCH.TALKIE: "talkie",
11391152
MODEL_ARCH.MELLUM: "mellum",
11401153
MODEL_ARCH.NANBEIGE: "nanbeige",
@@ -1210,6 +1223,13 @@ class MODEL_TENSOR(IntEnum):
12101223
MODEL_TENSOR.FFN_GATE_TID2EID: "blk.{bid}.ffn_gate_tid2eid",
12111224
MODEL_TENSOR.MOE_LATENT_DOWN: "blk.{bid}.ffn_latent_down", # nemotron 3 super
12121225
MODEL_TENSOR.MOE_LATENT_UP: "blk.{bid}.ffn_latent_up", # nemotron 3 super
1226+
MODEL_TENSOR.MOE_LATENT_NORM: "blk.{bid}.ffn_latent_norm", # kimi k3
1227+
MODEL_TENSOR.ATTN_RES_NORM: "blk.{bid}.attn_res_norm", # kimi k3
1228+
MODEL_TENSOR.ATTN_RES_PROJ: "blk.{bid}.attn_res_proj", # kimi k3
1229+
MODEL_TENSOR.FFN_RES_NORM: "blk.{bid}.ffn_res_norm", # kimi k3
1230+
MODEL_TENSOR.FFN_RES_PROJ: "blk.{bid}.ffn_res_proj", # kimi k3
1231+
MODEL_TENSOR.OUTPUT_RES_NORM: "output_res_norm", # kimi k3
1232+
MODEL_TENSOR.OUTPUT_RES_PROJ: "output_res_proj", # kimi k3
12131233
MODEL_TENSOR.LAYER_OUT_NORM: "blk.{bid}.layer_output_norm",
12141234
MODEL_TENSOR.LAYER_OUT_SCALE: "blk.{bid}.layer_output_scale",
12151235
MODEL_TENSOR.PER_LAYER_TOKEN_EMBD: "per_layer_token_embd", # gemma3n
@@ -4438,6 +4458,55 @@ class MODEL_TENSOR(IntEnum):
44384458
MODEL_TENSOR.FFN_DOWN,
44394459
MODEL_TENSOR.FFN_UP,
44404460
],
4461+
MODEL_ARCH.KIMI_K3: [
4462+
MODEL_TENSOR.TOKEN_EMBD,
4463+
MODEL_TENSOR.OUTPUT_NORM,
4464+
MODEL_TENSOR.OUTPUT,
4465+
MODEL_TENSOR.OUTPUT_RES_NORM,
4466+
MODEL_TENSOR.OUTPUT_RES_PROJ,
4467+
MODEL_TENSOR.ATTN_NORM,
4468+
MODEL_TENSOR.ATTN_Q,
4469+
MODEL_TENSOR.ATTN_K,
4470+
MODEL_TENSOR.ATTN_V,
4471+
MODEL_TENSOR.ATTN_OUT,
4472+
MODEL_TENSOR.ATTN_GATE,
4473+
MODEL_TENSOR.ATTN_Q_A,
4474+
MODEL_TENSOR.ATTN_Q_B,
4475+
MODEL_TENSOR.ATTN_KV_A_MQA,
4476+
MODEL_TENSOR.ATTN_KV_B,
4477+
MODEL_TENSOR.ATTN_K_B,
4478+
MODEL_TENSOR.ATTN_V_B,
4479+
MODEL_TENSOR.ATTN_Q_A_NORM,
4480+
MODEL_TENSOR.ATTN_KV_A_NORM,
4481+
MODEL_TENSOR.ATTN_RES_NORM,
4482+
MODEL_TENSOR.ATTN_RES_PROJ,
4483+
MODEL_TENSOR.FFN_RES_NORM,
4484+
MODEL_TENSOR.FFN_RES_PROJ,
4485+
MODEL_TENSOR.FFN_NORM,
4486+
MODEL_TENSOR.FFN_GATE,
4487+
MODEL_TENSOR.FFN_DOWN,
4488+
MODEL_TENSOR.FFN_UP,
4489+
MODEL_TENSOR.FFN_GATE_INP,
4490+
MODEL_TENSOR.FFN_GATE_EXP,
4491+
MODEL_TENSOR.FFN_DOWN_EXP,
4492+
MODEL_TENSOR.FFN_UP_EXP,
4493+
MODEL_TENSOR.MOE_LATENT_DOWN,
4494+
MODEL_TENSOR.MOE_LATENT_NORM,
4495+
MODEL_TENSOR.MOE_LATENT_UP,
4496+
MODEL_TENSOR.SSM_CONV1D_Q,
4497+
MODEL_TENSOR.SSM_CONV1D_K,
4498+
MODEL_TENSOR.SSM_CONV1D_V,
4499+
MODEL_TENSOR.SSM_F_A,
4500+
MODEL_TENSOR.SSM_F_B,
4501+
MODEL_TENSOR.SSM_BETA,
4502+
MODEL_TENSOR.SSM_A,
4503+
MODEL_TENSOR.SSM_DT,
4504+
MODEL_TENSOR.SSM_NORM,
4505+
MODEL_TENSOR.FFN_EXP_PROBS_B,
4506+
MODEL_TENSOR.FFN_GATE_SHEXP,
4507+
MODEL_TENSOR.FFN_DOWN_SHEXP,
4508+
MODEL_TENSOR.FFN_UP_SHEXP,
4509+
],
44414510
MODEL_ARCH.KIMI_LINEAR: [
44424511
MODEL_TENSOR.TOKEN_EMBD,
44434512
MODEL_TENSOR.OUTPUT_NORM,

gguf-py/gguf/gguf_writer.py

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -881,6 +881,15 @@ def add_moe_every_n_layers(self, value: int) -> None:
881881
def add_moe_latent_size(self, value: int) -> None:
882882
self.add_uint32(Keys.LLM.MOE_LATENT_SIZE.format(arch=self.arch), value)
883883

884+
def add_situ_beta(self, value: float) -> None:
885+
self.add_float32(Keys.LLM.SITU_BETA.format(arch=self.arch), value)
886+
887+
def add_situ_linear_beta(self, value: float) -> None:
888+
self.add_float32(Keys.LLM.SITU_LINEAR_BETA.format(arch=self.arch), value)
889+
890+
def add_attn_res_block_size(self, value: int) -> None:
891+
self.add_uint32(Keys.LLM.ATTN_RES_BLOCK_SIZE.format(arch=self.arch), value)
892+
884893
def add_nextn_predict_layers(self, count: int) -> None:
885894
self.add_uint32(Keys.LLM.NEXTN_PREDICT_LAYERS.format(arch=self.arch), count)
886895

@@ -1084,6 +1093,9 @@ def add_ssm_dt_b_c_rms(self, value: bool) -> None:
10841093
def add_kda_head_dim(self, value: int) -> None:
10851094
self.add_uint32(Keys.KDA.HEAD_DIM.format(arch=self.arch), value)
10861095

1096+
def add_kda_gate_lower_bound(self, value: float) -> None:
1097+
self.add_float32(Keys.KDA.GATE_LOWER_BOUND.format(arch=self.arch), value)
1098+
10871099
def add_tokenizer_model(self, model: str) -> None:
10881100
self.add_string(Keys.Tokenizer.MODEL, model)
10891101

gguf-py/gguf/tensor_mapping.py

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -118,6 +118,14 @@ class TensorNameMap:
118118
"model.norm", # cogvlm
119119
),
120120

121+
MODEL_TENSOR.OUTPUT_RES_NORM: (
122+
"model.output_attn_res_norm", # kimi k3
123+
),
124+
125+
MODEL_TENSOR.OUTPUT_RES_PROJ: (
126+
"model.output_attn_res_proj", # kimi k3
127+
),
128+
121129
# Rope frequencies
122130
MODEL_TENSOR.ROPE_FREQS: (
123131
"rope.freqs", # llama-pth
@@ -611,10 +619,32 @@ class TensorNameMap:
611619

612620
MODEL_TENSOR.MOE_LATENT_DOWN: (
613621
"backbone.layers.{bid}.mixer.fc1_latent_proj", # nemotron 3 super
622+
"model.layers.{bid}.block_sparse_moe.routed_expert_down_proj", # kimi k3
614623
),
615624

616625
MODEL_TENSOR.MOE_LATENT_UP: (
617626
"backbone.layers.{bid}.mixer.fc2_latent_proj", # nemotron 3 super
627+
"model.layers.{bid}.block_sparse_moe.routed_expert_up_proj", # kimi k3
628+
),
629+
630+
MODEL_TENSOR.MOE_LATENT_NORM: (
631+
"model.layers.{bid}.block_sparse_moe.routed_expert_norm", # kimi k3
632+
),
633+
634+
MODEL_TENSOR.ATTN_RES_NORM: (
635+
"model.layers.{bid}.self_attention_res_norm", # kimi k3
636+
),
637+
638+
MODEL_TENSOR.ATTN_RES_PROJ: (
639+
"model.layers.{bid}.self_attention_res_proj", # kimi k3
640+
),
641+
642+
MODEL_TENSOR.FFN_RES_NORM: (
643+
"model.layers.{bid}.mlp_res_norm", # kimi k3
644+
),
645+
646+
MODEL_TENSOR.FFN_RES_PROJ: (
647+
"model.layers.{bid}.mlp_res_proj", # kimi k3
618648
),
619649

620650
# Feed-forward down

0 commit comments

Comments
 (0)