Files
K3/kda/models/k3_config.py
T
dela d1da0816f2 LatentMoE: sparse permute-dispatch + padded bmm
Replace dense all-expert forward (16 experts × all tokens) with
permute-dispatch: sort token-expert pairs by expert id, pad to
[R, C, ℓ] (C = max tokens per expert), run 3 bmm calls for the
batched SiTU-GLU activation, then scatter-add weighted results back.

Routed expert FLOPs drop from R·N to R·C (C ≈ N·k/R under uniform
routing). SiTU parameter structure unchanged; checkpoint compatible.

Tests: sparse-vs-dense fwd/bwd equivalence, unselected expert zero
grad, last_capacity tracking.
2026-08-25 17:47:44 +08:00

129 lines
4.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""K3Config — Kimi K3 架构的小规模复现配置 (KDA + Gated MLA + Stable LatentMoE).
对照 learning/kimi-k3-notes §尺寸速查 (真实 K3 → 本 toy 缩比):
hidden 7168 → 256; L 93 → 4; H=HV 96 → 8; K=V 128 → 16;
MLA kv_lora 512 → 32, q_lora 1536 → 64, nope/v 128 → 16;
MoE ℓ=d/2=3584 → 128, 896/16 → 16/2, shared 2, d_ff 3072 → 96.
Hybrid Attention (K3): 每 4 层 1 次 Gated MLA, 末层强制 MLA.
Presets:
toy — ~8M, 自训 8k SP, 本地过拟合
0.5b — ~482M, Qwen3 词表, 32–40GB bf16;默认 step 是冒烟,翻译前置用 --max-tokens
"""
from __future__ import annotations
from dataclasses import dataclass
# Qwen3 config.json; train_k3 overrides with len(tokenizer).
QWEN3_VOCAB_SIZE = 151936
@dataclass
class K3Config:
# 主干
hidden_size: int = 256
num_hidden_layers: int = 4
vocab_size: int = 8192 # toy: data/spm_4k; 0.5b: Qwen3
initializer_range: float = 0.02
norm_eps: float = 1e-6
tie_word_embeddings: bool = False
max_position_embeddings: int = 2048 # NoPE, 仅语义保留
# KDA (K3: H = HV = 96, 无 GVA)
num_heads: int = 8
head_dim: int = 16
chunk_size: int = 16
lower_bound: float | None = -5.0
use_gate_in_kernel: bool = True
use_qk_l2norm_in_kernel: bool = True
use_beta_sigmoid_in_kernel: bool = True
# Gated MLA (NoPE)
kv_lora_rank: int = 32
q_lora_rank: int = 64
qk_nope_head_dim: int = 16
v_head_dim: int = 16
# Stable LatentMoE
moe_latent_size: int = 128 # ℓ = d/2
n_routed: int = 16
top_k: int = 2
n_shared: int = 2
moe_d_ff: int = 96
situ_beta1: float = 4.0
situ_beta2: float = 25.0
kda_backend: str = "reference"
# Depth mixer. off = DecoderBlock residual; block matches K3.
attnres: str = "off" # off | full | block
attnres_block_size: int | None = None # DecoderBlocks / AttnRes block; None ≈ L/8
attnres_zero_init_queries: bool = True
attnres_final_aggregate: bool = True
gradient_checkpointing: bool = False
def __post_init__(self):
from ..layers.attn_res import validate_attnres
validate_attnres(self.attnres, self.attnres_block_size)
@classmethod
def preset(cls, name: str) -> K3Config:
if name == "toy":
return cls()
if name in {"0.5b", "500m"}:
# H * head_dim == hidden. Routed 16 Top-2; LatentMoE padded bmm.
# ~482M with tied Qwen3 embeddings. 6×(3 KDA + 1 MLA).
return cls(
hidden_size=768,
num_hidden_layers=24,
vocab_size=QWEN3_VOCAB_SIZE,
tie_word_embeddings=True,
max_position_embeddings=2048,
num_heads=12,
head_dim=64,
chunk_size=64,
kv_lora_rank=192,
q_lora_rank=512,
qk_nope_head_dim=64,
v_head_dim=64,
moe_latent_size=384,
n_routed=16,
top_k=2,
n_shared=2,
moe_d_ff=512,
# The pure-PyTorch reference is far too slow at this size.
kda_backend="triton",
gradient_checkpointing=True,
)
raise ValueError(f"unknown preset: {name}")
@property
def H(self) -> int:
return self.num_heads
@property
def HV(self) -> int:
return self.num_heads
@property
def K(self) -> int:
return self.head_dim
@property
def V(self) -> int:
return self.head_dim
def layer_types(self) -> list[str]:
"""Hybrid pattern: 每 4 层 1 次 MLA (0-based 层 3,7,...), 末层强制 MLA."""
types = ["kda"] * self.num_hidden_layers
for i in range(self.num_hidden_layers):
if i % 4 == 3:
types[i] = "mla"
types[-1] = "mla"
return types
def layer_specs(self) -> list[tuple[str, str]]:
return [(kind, "moe") for kind in self.layer_types()]