"""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: LatentMoE still runs every expert. # ~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()]