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.
129 lines
4.0 KiB
Python
129 lines
4.0 KiB
Python
"""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()]
|