Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
This commit is contained in:
@@ -0,0 +1,7 @@
|
||||
"""Configs and the single CausalLM entry."""
|
||||
|
||||
from .causal_lm import CausalLM
|
||||
from .config import KDAConfig
|
||||
from .k3_config import K3Config
|
||||
|
||||
__all__ = ["CausalLM", "K3Config", "KDAConfig"]
|
||||
@@ -0,0 +1,123 @@
|
||||
"""Causal LM stem: embed -> DecoderBlock* -> norm -> lm_head.
|
||||
|
||||
KDA-only and K3-like both use this class. Config.layer_specs() chooses
|
||||
attn/ffn per layer: ("kda"|"mla", "swiglu"|"moe").
|
||||
|
||||
``config.attnres`` selects the depth mixer:
|
||||
off — standard residual inside each DecoderBlock (default)
|
||||
full — Full AttnRes over attn|ffn sublayers
|
||||
block — Block AttnRes (K3); block size from ``attnres_block_size``
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from torch.utils.checkpoint import checkpoint as activation_checkpoint
|
||||
|
||||
from ..layers.attn_res import (
|
||||
BlockAttnResStack,
|
||||
BorrowedSubLayer,
|
||||
FullAttnResStack,
|
||||
atomic_block_size,
|
||||
)
|
||||
from ..layers.block import DecoderBlock
|
||||
from ..layers.rmsnorm import RMSNorm
|
||||
|
||||
|
||||
def _build_mixer(config, blocks: nn.ModuleList):
|
||||
mode = getattr(config, "attnres", "off")
|
||||
if mode == "off":
|
||||
return None
|
||||
atomics = []
|
||||
for block in blocks:
|
||||
atomics.append(BorrowedSubLayer(block.attn_norm, block.attn))
|
||||
atomics.append(BorrowedSubLayer(block.ffn_norm, block.ffn))
|
||||
kwargs = dict(
|
||||
eps=config.norm_eps,
|
||||
zero_init_queries=getattr(config, "attnres_zero_init_queries", True),
|
||||
is_final_aggregate=getattr(config, "attnres_final_aggregate", True),
|
||||
)
|
||||
if mode == "full":
|
||||
return FullAttnResStack(config.hidden_size, atomics, **kwargs)
|
||||
if mode == "block":
|
||||
return BlockAttnResStack(
|
||||
config.hidden_size,
|
||||
atomics,
|
||||
block_size=atomic_block_size(
|
||||
config.num_hidden_layers, getattr(config, "attnres_block_size", None)
|
||||
),
|
||||
**kwargs,
|
||||
)
|
||||
raise ValueError(f"unknown attnres mode: {mode!r}")
|
||||
|
||||
|
||||
class CausalLM(nn.Module):
|
||||
def __init__(self, config):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.attnres = getattr(config, "attnres", "off")
|
||||
self.embedding = nn.Embedding(config.vocab_size, config.hidden_size)
|
||||
self.blocks = nn.ModuleList(
|
||||
[
|
||||
DecoderBlock.from_spec(config, attn, ffn)
|
||||
for attn, ffn in config.layer_specs()
|
||||
]
|
||||
)
|
||||
self.mixer = _build_mixer(config, self.blocks)
|
||||
self.gradient_checkpointing = bool(
|
||||
getattr(config, "gradient_checkpointing", False)
|
||||
)
|
||||
self.norm = RMSNorm(config.hidden_size, config.norm_eps)
|
||||
self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
|
||||
nn.init.normal_(self.embedding.weight, std=config.initializer_range)
|
||||
nn.init.normal_(self.lm_head.weight, std=config.initializer_range)
|
||||
if config.tie_word_embeddings:
|
||||
self.lm_head.weight = self.embedding.weight
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
labels: torch.Tensor | None = None,
|
||||
ignore_index: int = -100,
|
||||
):
|
||||
x = self.embedding(input_ids)
|
||||
if self.mixer is None:
|
||||
for block in self.blocks:
|
||||
if self.gradient_checkpointing and self.training:
|
||||
x = activation_checkpoint(block, x, use_reentrant=False)
|
||||
else:
|
||||
x = block(x)
|
||||
elif self.gradient_checkpointing and self.training:
|
||||
x = activation_checkpoint(self.mixer, x, use_reentrant=False)
|
||||
else:
|
||||
x = self.mixer(x)
|
||||
logits = self.lm_head(self.norm(x))
|
||||
if labels is None:
|
||||
return logits
|
||||
return F.cross_entropy(
|
||||
logits[:, :-1].reshape(-1, logits.size(-1)),
|
||||
labels[:, 1:].reshape(-1),
|
||||
ignore_index=ignore_index,
|
||||
)
|
||||
|
||||
@torch.inference_mode()
|
||||
def generate(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
max_new_tokens: int,
|
||||
temperature: float = 0.0,
|
||||
eos_token_id: int | None = None,
|
||||
):
|
||||
for _ in range(max_new_tokens):
|
||||
logits = self(input_ids)[:, -1]
|
||||
if temperature > 0:
|
||||
probs = F.softmax(logits / temperature, dim=-1)
|
||||
next_token = torch.multinomial(probs, 1)
|
||||
else:
|
||||
next_token = logits.argmax(-1, keepdim=True)
|
||||
input_ids = torch.cat((input_ids, next_token), dim=1)
|
||||
if eos_token_id is not None and (next_token.squeeze(-1) == eos_token_id).all():
|
||||
break
|
||||
return input_ids
|
||||
@@ -0,0 +1,62 @@
|
||||
"""KDAConfig — toy Causal LM hyperparameters.
|
||||
|
||||
Defaults match the working reference-backend model: GVA with G=2,
|
||||
safe gate (lower_bound=-5), q/k L2-norm and beta sigmoid inside the op.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass
|
||||
class KDAConfig:
|
||||
hidden_size: int = 64
|
||||
num_hidden_layers: int = 2
|
||||
num_heads: int = 4
|
||||
num_value_heads: int = 8 # G = num_value_heads // num_heads
|
||||
head_dim: int = 16
|
||||
chunk_size: int = 16
|
||||
vocab_size: int = 256
|
||||
intermediate_size: int = 128
|
||||
max_position_embeddings: int = 128
|
||||
initializer_range: float = 0.02
|
||||
norm_eps: float = 1e-6
|
||||
use_gate_in_kernel: bool = True
|
||||
use_qk_l2norm_in_kernel: bool = True
|
||||
use_beta_sigmoid_in_kernel: bool = True
|
||||
lower_bound: float | None = -5.0
|
||||
tie_word_embeddings: bool = False
|
||||
kda_backend: str = "reference" # reference | triton | fla
|
||||
attnres: str = "off" # off | full | block
|
||||
attnres_block_size: int | None = None # DecoderBlocks / block; None ≈ L/8
|
||||
attnres_zero_init_queries: bool = True
|
||||
attnres_final_aggregate: bool = True
|
||||
gradient_checkpointing: bool = False
|
||||
|
||||
@property
|
||||
def H(self) -> int: return self.num_heads
|
||||
|
||||
@property
|
||||
def G(self) -> int: return self.num_value_heads // self.num_heads
|
||||
|
||||
@property
|
||||
def HV(self) -> int: return self.num_value_heads
|
||||
|
||||
@property
|
||||
def K(self) -> int: return self.head_dim
|
||||
|
||||
@property
|
||||
def V(self) -> int: return self.head_dim
|
||||
|
||||
def __post_init__(self):
|
||||
from ..layers.attn_res import validate_attnres
|
||||
|
||||
if self.num_value_heads % self.num_heads:
|
||||
raise ValueError("num_value_heads must be divisible by num_heads")
|
||||
supported = {"reference", "triton", "fla", "torch", "auto"}
|
||||
if self.kda_backend not in supported:
|
||||
raise ValueError(f"kda_backend must be one of {sorted(supported)}")
|
||||
validate_attnres(self.attnres, self.attnres_block_size)
|
||||
|
||||
def layer_specs(self) -> list[tuple[str, str]]:
|
||||
return [("kda", "swiglu")] * self.num_hidden_layers
|
||||
@@ -0,0 +1,128 @@
|
||||
"""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()]
|
||||
Reference in New Issue
Block a user