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:
dela
2026-08-25 14:43:17 +08:00
commit 584f7e9e73
140 changed files with 21592 additions and 0 deletions
+123
View File
@@ -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