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,23 @@
|
||||
"""Composable mixing layers: attn and ffn both map [B,T,D] -> [B,T,D].
|
||||
|
||||
Depth mixing (AttnRes) is not a layer_specs kind. CausalLM reads
|
||||
``config.attnres`` (off | full | block) and wraps DecoderBlock sublayers.
|
||||
"""
|
||||
|
||||
from .block import DecoderBlock, build_attn, build_ffn
|
||||
from .kda_attn import KDAAttention
|
||||
from .latent_moe import LatentMoE
|
||||
from .mla import GatedMLA
|
||||
from .rmsnorm import RMSNorm
|
||||
from .swiglu import SwiGLUMLP
|
||||
|
||||
__all__ = [
|
||||
"DecoderBlock",
|
||||
"GatedMLA",
|
||||
"KDAAttention",
|
||||
"LatentMoE",
|
||||
"RMSNorm",
|
||||
"SwiGLUMLP",
|
||||
"build_attn",
|
||||
"build_ffn",
|
||||
]
|
||||
@@ -0,0 +1,519 @@
|
||||
"""
|
||||
Attention Residual in one file
|
||||
|
||||
Reference:
|
||||
Kimi Team, Guangyu Chen, Yu Zhang, Jianlin Su, Weixin Xu, Siyuan Pan,
|
||||
Yaoyu Wang, Yucheng Wang, Guanduo Chen, et al.
|
||||
"Attention Residuals." arXiv:2603.15031, 2026.
|
||||
https://arxiv.org/abs/2603.15031
|
||||
|
||||
This module is a compact PyTorch reference implementation of:
|
||||
- Full AttnRes
|
||||
- Block AttnRes
|
||||
- two-phase inter/intra-block computation from the paper
|
||||
|
||||
CausalLM wires Full/Block stacks when ``config.attnres`` is ``full`` or
|
||||
``block``. Standard residual (``x += attn; x += ffn``) is ``attnres="off"``.
|
||||
"""
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from torch import Tensor, nn
|
||||
|
||||
|
||||
ATTNRES_MODES = ("off", "full", "block")
|
||||
|
||||
|
||||
def exists(x):
|
||||
return x is not None
|
||||
|
||||
|
||||
def validate_attnres(mode: str, block_size: int | None) -> None:
|
||||
if mode not in ATTNRES_MODES:
|
||||
raise ValueError(f"attnres must be one of {ATTNRES_MODES}, got {mode!r}")
|
||||
if block_size is not None and block_size < 1:
|
||||
raise ValueError(f"attnres_block_size must be >= 1, got {block_size}")
|
||||
|
||||
|
||||
def atomic_block_size(num_hidden_layers: int, attnres_block_size: int | None) -> int:
|
||||
"""DecoderBlocks per AttnRes block, converted to attn|ffn atomic layers.
|
||||
|
||||
``None`` targets about 8 blocks: ``max(1, ceil(L / 8))`` DecoderBlocks.
|
||||
"""
|
||||
layers_per_block = (
|
||||
attnres_block_size
|
||||
if attnres_block_size is not None
|
||||
else max(1, (num_hidden_layers + 7) // 8)
|
||||
)
|
||||
if layers_per_block < 1:
|
||||
raise ValueError(f"attnres_block_size must be >= 1, got {layers_per_block}")
|
||||
return layers_per_block * 2
|
||||
|
||||
|
||||
class BorrowedSubLayer(nn.Module):
|
||||
"""``fn(norm(x))`` without registering ``norm``/``fn`` (owned by DecoderBlock)."""
|
||||
|
||||
def __init__(self, norm: nn.Module, fn: nn.Module):
|
||||
super().__init__()
|
||||
self._borrowed = (norm, fn)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
norm, fn = self._borrowed
|
||||
return fn(norm(x))
|
||||
|
||||
|
||||
def rms(x: Tensor, eps: float):
|
||||
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + eps)
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
def __init__(self, dim: int, eps: float):
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return rms(x, self.eps) * self.weight
|
||||
|
||||
|
||||
class DepthResidual(nn.Module):
|
||||
"""
|
||||
h_l = sum_i softmax_i(w_l^T RMSNorm(v_i))*v_i
|
||||
|
||||
Keep query and RMSNorm gain separate
|
||||
Since q^T (gamma * RMS(v)) == (q * gamma)^T RMS(v),
|
||||
we can fold gamma into q for scoring.
|
||||
"""
|
||||
|
||||
def __init__(self, dim: int, eps: float = 1e-8, zero_init: bool = True):
|
||||
super().__init__()
|
||||
self.query = nn.Parameter(torch.zeros(dim))
|
||||
self.norm = RMSNorm(dim, eps=eps)
|
||||
|
||||
if not zero_init:
|
||||
nn.init.normal_(self.query, std=0.02)
|
||||
|
||||
def effective_query(self) -> Tensor:
|
||||
return (self.query * self.norm.weight).float()
|
||||
|
||||
def logits(self, sources: Tensor | list[Tensor] | tuple[Tensor, ...]) -> Tensor:
|
||||
sources = stack_layers(sources) # [n, b, t, d]
|
||||
q = self.effective_query() # [d]
|
||||
k = rms(sources.float(), self.norm.eps) # [n, b, t, d]
|
||||
return torch.einsum("d, n b t d -> n b t", q, k)
|
||||
|
||||
def forward(self, sources: Tensor | list[Tensor] | tuple[Tensor, ...]) -> Tensor:
|
||||
sources = stack_layers(sources)
|
||||
weights = self.logits(sources).softmax(dim=0)
|
||||
out = torch.einsum("n b t, n b t d -> b t d", weights, sources.float())
|
||||
return out.to(sources.dtype)
|
||||
|
||||
|
||||
class DepthResidualList(nn.Module):
|
||||
def __init__(self, dim: int, depth: int, eps: float, zero_init: bool = True):
|
||||
super().__init__()
|
||||
# for L layers (depth), create depth residual modules
|
||||
self.layers = nn.ModuleList(
|
||||
[DepthResidual(dim, eps=eps, zero_init=zero_init) for _ in range(depth)]
|
||||
)
|
||||
|
||||
def __getitem__(self, idx: int) -> DepthResidual:
|
||||
return self.layers[idx]
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self.layers)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.layers)
|
||||
|
||||
|
||||
# attnres stacks
|
||||
|
||||
|
||||
class FullAttnResStack(nn.Module):
|
||||
"""
|
||||
Full AttnRes over atomic layers
|
||||
eg: f_1,...,f_L
|
||||
Each entry in `layers` should already be a full atomic layer fxn
|
||||
x -> f_l(x)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
layers,
|
||||
*,
|
||||
eps: float = 1e-8,
|
||||
zero_init_queries: bool = True,
|
||||
is_final_aggregate: bool = True,
|
||||
):
|
||||
super().__init__()
|
||||
self.layers = nn.ModuleList(list(layers))
|
||||
self.eps = eps
|
||||
|
||||
depth = len(self.layers)
|
||||
self.residuals = DepthResidualList(dim, depth, eps, zero_init_queries)
|
||||
self.final_residual = (
|
||||
DepthResidual(dim, eps, zero_init_queries) if is_final_aggregate else None
|
||||
)
|
||||
|
||||
def forward_naive(self, x: Tensor) -> Tensor:
|
||||
sources = [x]
|
||||
for layer, residual in zip(self.layers, self.residuals):
|
||||
h = residual(sources)
|
||||
out = layer(h)
|
||||
sources.append(out)
|
||||
|
||||
return (
|
||||
self.final_residual(sources) if exists(self.final_residual) else sources[-1]
|
||||
)
|
||||
|
||||
def forward_two_phase(self, x: Tensor, schedule_block_size: int) -> Tensor:
|
||||
assert schedule_block_size > 0
|
||||
sources = [x]
|
||||
depth = len(self.layers)
|
||||
|
||||
start = 0
|
||||
while start < depth:
|
||||
end = min(start + schedule_block_size, depth)
|
||||
queries = torch.stack(
|
||||
[self.residuals[i].effective_query() for i in range(start, end)], dim=0
|
||||
)
|
||||
inter_sources = stack_layers(sources)
|
||||
inter_stats = attn_with_stats(queries, inter_sources, self.eps)
|
||||
|
||||
local_outputs = [] # outputs of intra-block
|
||||
for local_idx, layer_idx in enumerate(range(start, end)):
|
||||
stats = inter_stats.select(local_idx)
|
||||
if len(local_outputs) > 0:
|
||||
intra_sources = stack_layers(local_outputs)
|
||||
intra = attn_with_stats(
|
||||
queries[local_idx : local_idx + 1], intra_sources, self.eps
|
||||
).select(0)
|
||||
stats = merge_attn_stats(stats, intra)
|
||||
h = stats.normalized()
|
||||
out = self.layers[layer_idx](h)
|
||||
local_outputs.append(out)
|
||||
sources.append(out)
|
||||
|
||||
start = end
|
||||
|
||||
return (
|
||||
self.final_residual(sources) if exists(self.final_residual) else sources[-1]
|
||||
)
|
||||
|
||||
def forward(self, x: Tensor, schedule_block_size: int | None = None) -> Tensor:
|
||||
if schedule_block_size is None:
|
||||
return self.forward_naive(x)
|
||||
return self.forward_two_phase(x, schedule_block_size)
|
||||
|
||||
|
||||
class BlockAttnResStack(nn.Module):
|
||||
"""
|
||||
Block AttnRes over atomic layers
|
||||
|
||||
`block_size` is in atomic layers, not Transformer blocks.
|
||||
Eg: block_size=8 -> 4 transformer blocks when layers alternate attn/MLP
|
||||
|
||||
The default forward path is the two-phase algorithm from the paper:
|
||||
phase 1: batch inter-block attn from all queries in the block
|
||||
phase 2: merge the evolving intra-block partial sum with online softmax
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
layers,
|
||||
*,
|
||||
block_size: int,
|
||||
eps: float = 1e-8,
|
||||
zero_init_queries: bool = True,
|
||||
is_final_aggregate: bool = True,
|
||||
):
|
||||
super().__init__()
|
||||
self.layers = nn.ModuleList(list(layers))
|
||||
assert len(self.layers) > 0
|
||||
assert block_size > 0
|
||||
self.block_size = block_size
|
||||
self.eps = eps
|
||||
|
||||
depth = len(self.layers)
|
||||
self.residuals = DepthResidualList(
|
||||
dim, depth, eps=eps, zero_init=zero_init_queries
|
||||
)
|
||||
self.final_residual = (
|
||||
DepthResidual(dim, eps=eps, zero_init=zero_init_queries)
|
||||
if is_final_aggregate
|
||||
else None
|
||||
)
|
||||
|
||||
def forward_naive(self, x: Tensor) -> Tensor:
|
||||
blocks = [x] # b_0=embedding/input representation
|
||||
partial = None
|
||||
|
||||
for layer_idx, (layer, residual) in enumerate(
|
||||
zip(self.layers, self.residuals), start=1
|
||||
):
|
||||
sources = blocks if partial is None else blocks + [partial]
|
||||
h = residual(sources)
|
||||
out = layer(h)
|
||||
partial = out if partial is None else (partial + out)
|
||||
|
||||
if (layer_idx % self.block_size == 0) or (layer_idx == len(self.layers)):
|
||||
blocks.append(partial)
|
||||
partial = None
|
||||
|
||||
return (
|
||||
self.final_residual(blocks) if exists(self.final_residual) else blocks[-1]
|
||||
)
|
||||
|
||||
def _run_block_two_phase(
|
||||
self, blocks: list[Tensor], start: int, end: int
|
||||
) -> Tensor:
|
||||
queries = torch.stack(
|
||||
[self.residuals[i].effective_query() for i in range(start, end)], dim=0
|
||||
)
|
||||
inter_sources = stack_layers(blocks)
|
||||
inter = attn_with_stats(queries, inter_sources, self.eps)
|
||||
partial = None
|
||||
for local_idx, layer_idx in enumerate(range(start, end)):
|
||||
stats = inter.select(local_idx)
|
||||
|
||||
if partial is not None:
|
||||
intra = single_source_stats(queries[local_idx], partial, self.eps)
|
||||
stats = merge_attn_stats(stats, intra)
|
||||
|
||||
h = stats.normalized()
|
||||
out = self.layers[layer_idx](h)
|
||||
partial = out if partial is None else (partial + out)
|
||||
|
||||
return partial
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
blocks = [x]
|
||||
depth = len(self.layers)
|
||||
start = 0
|
||||
while start < depth:
|
||||
end = min(start + self.block_size, depth)
|
||||
blocks.append(self._run_block_two_phase(blocks, start, end))
|
||||
start = end
|
||||
|
||||
return (
|
||||
self.final_residual(blocks) if exists(self.final_residual) else blocks[-1]
|
||||
)
|
||||
|
||||
|
||||
# helpers
|
||||
|
||||
|
||||
def stack_layers(sources: Tensor | list[Tensor] | tuple[Tensor, ...]) -> Tensor:
|
||||
if isinstance(sources, Tensor):
|
||||
assert sources.ndim == 4, f"expected [n, b, t, d] got {tuple(sources.shape)}"
|
||||
return sources
|
||||
assert len(sources) > 0, "needs at least one source"
|
||||
return torch.stack(tuple(sources), dim=0)
|
||||
|
||||
|
||||
class SingleAttnStats:
|
||||
def __init__(self, numer: Tensor, denom: Tensor, max: Tensor):
|
||||
self.numer = numer # [b,t,d]
|
||||
self.max = max # [b,t]
|
||||
self.denom = denom # [b,t]
|
||||
|
||||
def normalized(self) -> Tensor:
|
||||
return self.numer / self.denom[..., None]
|
||||
|
||||
|
||||
class AttnStats:
|
||||
# store the numerator => e^{s_{j}-m} * v_j where m is the max score so far
|
||||
# store the max m = max(s_j)
|
||||
# store the denominator sum_j e^{s_{j}-m}
|
||||
def __init__(self, numer: Tensor, denom: Tensor, max: Tensor):
|
||||
self.numer = numer # [q,b,t,d]
|
||||
self.max = max # [q,b,t]
|
||||
self.denom = denom # [q,b,t]
|
||||
|
||||
def select(self, idx: int) -> "SingleAttnStats":
|
||||
return SingleAttnStats(self.numer[idx], self.denom[idx], self.max[idx])
|
||||
|
||||
|
||||
def attn_with_stats(queries: Tensor, sources: Tensor, eps: float = 1e-8) -> AttnStats:
|
||||
"""
|
||||
queries: [q, d]
|
||||
sources: [n, b, t, d]
|
||||
|
||||
Returns the following for online softmax:
|
||||
numer = sum_i exp(logit_i - m)*v_i
|
||||
m = max_i logit_i
|
||||
denom = sum_i exp(logit_i - m)
|
||||
"""
|
||||
normed = rms(sources, eps)
|
||||
logits = torch.einsum("q d, n b t d -> q n b t", queries, normed)
|
||||
m = logits.amax(dim=1)
|
||||
weights = torch.exp(logits - m[:, None])
|
||||
numer = torch.einsum("q n b t, n b t d -> q b t d", weights, sources)
|
||||
denom = weights.sum(dim=1)
|
||||
return AttnStats(numer, denom, m)
|
||||
|
||||
|
||||
def single_source_stats(
|
||||
query: Tensor, source: Tensor, eps: float = 1e-8
|
||||
) -> SingleAttnStats:
|
||||
score = torch.einsum("d, b t d -> b t", query, rms(source, eps))
|
||||
denom = torch.ones_like(score)
|
||||
return SingleAttnStats(source, denom, score)
|
||||
|
||||
|
||||
def merge_attn_stats(a: SingleAttnStats, b: SingleAttnStats) -> SingleAttnStats:
|
||||
m = torch.maximum(a.max, b.max)
|
||||
wa = torch.exp(a.max - m)
|
||||
wb = torch.exp(b.max - m)
|
||||
numer = wa[..., None] * a.numer + wb[..., None] * b.numer
|
||||
denom = wa * a.denom + wb * b.denom
|
||||
return SingleAttnStats(numer, denom, m)
|
||||
|
||||
|
||||
# transformer
|
||||
class PreNorm(nn.Module):
|
||||
def __init__(self, dim: int, fn: nn.Module, eps: float = 1e-8):
|
||||
super().__init__()
|
||||
self.norm = RMSNorm(dim, eps=eps)
|
||||
self.fn = fn
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
return self.fn(self.norm(x))
|
||||
|
||||
|
||||
class CausalAttention(nn.Module):
|
||||
def __init__(
|
||||
self, dim: int, heads: int = 8, dim_head: int = 64, dropout: float = 0.0
|
||||
):
|
||||
super().__init__()
|
||||
inner_dim = heads * dim_head
|
||||
self.heads = heads
|
||||
self.dim_head = dim_head
|
||||
self.dropout = dropout
|
||||
|
||||
self.to_qkv = nn.Linear(dim, inner_dim * 3, bias=False)
|
||||
self.to_out = nn.Linear(inner_dim, dim, bias=False)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
q, k, v = self.to_qkv(x).chunk(3, dim=-1)
|
||||
|
||||
def split_heads(y: Tensor) -> Tensor:
|
||||
return rearrange(y, "b t (h d) -> b h t d", h=self.heads)
|
||||
|
||||
q, k, v = map(split_heads, (q, k, v))
|
||||
|
||||
out = F.scaled_dot_product_attention(
|
||||
q, k, v, is_causal=True, dropout_p=self.dropout if self.training else 0.0
|
||||
)
|
||||
out = rearrange(out, "b h t d -> b t (h d)")
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class SwiGLU(nn.Module):
|
||||
def __init__(self, dim: int, mult: int = 4, dropout: float = 0.0):
|
||||
# dropout not needed unless training on a smaller training data
|
||||
super().__init__()
|
||||
inner_dim = dim * mult
|
||||
self.to_hidden = nn.Linear(dim, inner_dim * 2, bias=False)
|
||||
self.to_out = nn.Linear(inner_dim, dim, bias=False)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
|
||||
def forward(self, x: Tensor) -> Tensor:
|
||||
gate, value = self.to_hidden(x).chunk(2, dim=-1)
|
||||
x = F.silu(gate) * value
|
||||
x = self.dropout(x)
|
||||
return self.to_out(x)
|
||||
|
||||
|
||||
class AttnResTransformer(nn.Module):
|
||||
"""
|
||||
Small GPT-style reference model using AttnRes
|
||||
|
||||
Using plain PyTorch: tok/pos embedding, alternating
|
||||
causal attn, SwiGLU MLP layers, final norm, output head.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
num_tokens: int,
|
||||
dim: int,
|
||||
depth: int,
|
||||
max_seq_len: int,
|
||||
heads: int = 8,
|
||||
dim_head: int = 64,
|
||||
ff_mult: int = 4,
|
||||
attn_dropout: float = 0.0,
|
||||
ff_dropout: float = 0.0,
|
||||
attnres: str = "block", # full or block
|
||||
block_size: int = 8,
|
||||
zero_init_queries: bool = True,
|
||||
is_final_aggregate: bool = True,
|
||||
eps: float = 1e-8,
|
||||
):
|
||||
super().__init__()
|
||||
assert attnres in {"full", "block"}
|
||||
self.max_seq_len = max_seq_len
|
||||
self.attnres = attnres
|
||||
self.token_emb = nn.Embedding(num_tokens, dim)
|
||||
self.pos_emb = nn.Embedding(max_seq_len, dim)
|
||||
|
||||
atomic_layers = []
|
||||
for _ in range(depth):
|
||||
atomic_layers.append(
|
||||
PreNorm(dim, CausalAttention(dim, heads, dim_head, attn_dropout), eps)
|
||||
)
|
||||
atomic_layers.append(PreNorm(dim, SwiGLU(dim, ff_mult, ff_dropout), eps))
|
||||
if attnres == "full":
|
||||
self.backbone = FullAttnResStack(
|
||||
dim,
|
||||
atomic_layers,
|
||||
eps=eps,
|
||||
zero_init_queries=zero_init_queries,
|
||||
is_final_aggregate=is_final_aggregate,
|
||||
)
|
||||
else:
|
||||
self.backbone = BlockAttnResStack(
|
||||
dim,
|
||||
atomic_layers,
|
||||
block_size=block_size,
|
||||
eps=eps,
|
||||
zero_init_queries=zero_init_queries,
|
||||
is_final_aggregate=is_final_aggregate,
|
||||
)
|
||||
|
||||
self.final_norm = RMSNorm(dim, eps)
|
||||
self.to_logits = nn.Linear(dim, num_tokens, bias=False)
|
||||
|
||||
def forward(self, ids: Tensor, schedule_block_size: int | None = None) -> Tensor:
|
||||
b, t = ids.shape
|
||||
assert t <= self.max_seq_len
|
||||
pos = torch.arange(t, device=ids.device)
|
||||
x = self.token_emb(ids) + self.pos_emb(pos)[None, :, :]
|
||||
if self.attnres == "full":
|
||||
x = self.backbone(x, schedule_block_size=schedule_block_size)
|
||||
else:
|
||||
x = self.backbone(x)
|
||||
x = self.final_norm(x)
|
||||
return self.to_logits(x)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ATTNRES_MODES",
|
||||
"RMSNorm",
|
||||
"DepthResidual",
|
||||
"DepthResidualList",
|
||||
"FullAttnResStack",
|
||||
"BlockAttnResStack",
|
||||
"BorrowedSubLayer",
|
||||
"PreNorm",
|
||||
"CausalAttention",
|
||||
"SwiGLU",
|
||||
"AttnResTransformer",
|
||||
"atomic_block_size",
|
||||
"validate_attnres",
|
||||
]
|
||||
@@ -0,0 +1,51 @@
|
||||
"""Decoder block: x += attn(norm(x)); x += ffn(norm(x)).
|
||||
|
||||
attn/ffn are any modules with forward: [B,T,D] -> [B,T,D].
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from torch import nn
|
||||
|
||||
from .kda_attn import KDAAttention
|
||||
from .latent_moe import LatentMoE
|
||||
from .mla import GatedMLA
|
||||
from .rmsnorm import RMSNorm
|
||||
from .swiglu import SwiGLUMLP
|
||||
|
||||
|
||||
def build_attn(config, kind: str) -> nn.Module:
|
||||
if kind == "kda":
|
||||
return KDAAttention.from_config(config)
|
||||
if kind == "mla":
|
||||
return GatedMLA.from_config(config)
|
||||
raise ValueError(f"unknown attn kind: {kind}")
|
||||
|
||||
|
||||
def build_ffn(config, kind: str) -> nn.Module:
|
||||
if kind == "swiglu":
|
||||
return SwiGLUMLP.from_config(config)
|
||||
if kind == "moe":
|
||||
return LatentMoE.from_config(config)
|
||||
raise ValueError(f"unknown ffn kind: {kind}")
|
||||
|
||||
|
||||
class DecoderBlock(nn.Module):
|
||||
def __init__(self, hidden_size: int, norm_eps: float, attn: nn.Module, ffn: nn.Module):
|
||||
super().__init__()
|
||||
self.attn_norm = RMSNorm(hidden_size, norm_eps)
|
||||
self.attn = attn
|
||||
self.ffn_norm = RMSNorm(hidden_size, norm_eps)
|
||||
self.ffn = ffn
|
||||
|
||||
@classmethod
|
||||
def from_spec(cls, config, attn_kind: str, ffn_kind: str) -> DecoderBlock:
|
||||
return cls(
|
||||
config.hidden_size,
|
||||
config.norm_eps,
|
||||
build_attn(config, attn_kind),
|
||||
build_ffn(config, ffn_kind),
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = x + self.attn(self.attn_norm(x))
|
||||
return x + self.ffn(self.ffn_norm(x))
|
||||
@@ -0,0 +1,99 @@
|
||||
"""KDA attention: project q/k/v/g/beta, run chunk_kda, project back to D."""
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from ..ops.api import chunk_kda
|
||||
|
||||
|
||||
class KDAAttention(nn.Module):
|
||||
"""Mixing module: x [B,T,D] -> y [B,T,D]."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_heads: int,
|
||||
num_value_heads: int,
|
||||
head_dim: int,
|
||||
*,
|
||||
chunk_size: int = 16,
|
||||
initializer_range: float = 0.02,
|
||||
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,
|
||||
kda_backend: str = "reference",
|
||||
):
|
||||
super().__init__()
|
||||
if num_value_heads % num_heads:
|
||||
raise ValueError("num_value_heads must be divisible by num_heads")
|
||||
self.hidden_size = hidden_size
|
||||
self.num_heads = num_heads
|
||||
self.num_value_heads = num_value_heads
|
||||
self.head_dim = head_dim
|
||||
self.chunk_size = chunk_size
|
||||
self.initializer_range = initializer_range
|
||||
self.use_gate_in_kernel = use_gate_in_kernel
|
||||
self.use_qk_l2norm_in_kernel = use_qk_l2norm_in_kernel
|
||||
self.use_beta_sigmoid_in_kernel = use_beta_sigmoid_in_kernel
|
||||
self.lower_bound = lower_bound
|
||||
self.kda_backend = kda_backend
|
||||
|
||||
H, HV, K, V = num_heads, num_value_heads, head_dim, head_dim
|
||||
self.q_proj = nn.Linear(hidden_size, H * K, bias=False)
|
||||
self.k_proj = nn.Linear(hidden_size, H * K, bias=False)
|
||||
self.v_proj = nn.Linear(hidden_size, HV * V, bias=False)
|
||||
self.g_proj = nn.Linear(hidden_size, HV * K, bias=False)
|
||||
self.beta_proj = nn.Linear(hidden_size, HV, bias=False)
|
||||
self.o_proj = nn.Linear(HV * V, hidden_size, bias=False)
|
||||
self.A_log = nn.Parameter(torch.zeros(HV))
|
||||
# With safe_gate=-5, bias=-4 starts at g≈-0.09 (about 91% state retention).
|
||||
self.dt_bias = nn.Parameter(torch.full((HV, K), -4.0))
|
||||
self.apply(self._init_weights)
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config) -> KDAAttention:
|
||||
return cls(
|
||||
hidden_size=config.hidden_size,
|
||||
num_heads=config.num_heads,
|
||||
num_value_heads=getattr(config, "num_value_heads", config.num_heads),
|
||||
head_dim=config.head_dim,
|
||||
chunk_size=config.chunk_size,
|
||||
initializer_range=config.initializer_range,
|
||||
use_gate_in_kernel=config.use_gate_in_kernel,
|
||||
use_qk_l2norm_in_kernel=config.use_qk_l2norm_in_kernel,
|
||||
use_beta_sigmoid_in_kernel=config.use_beta_sigmoid_in_kernel,
|
||||
lower_bound=config.lower_bound,
|
||||
kda_backend=config.kda_backend,
|
||||
)
|
||||
|
||||
def _init_weights(self, module):
|
||||
if isinstance(module, nn.Linear):
|
||||
nn.init.normal_(module.weight, std=self.initializer_range)
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
B, T, _ = x.shape
|
||||
H, HV, K, V = self.num_heads, self.num_value_heads, self.head_dim, self.head_dim
|
||||
q = self.q_proj(x).view(B, T, H, K)
|
||||
k = self.k_proj(x).view(B, T, H, K)
|
||||
v = self.v_proj(x).view(B, T, HV, V)
|
||||
g_raw = self.g_proj(x).view(B, T, HV, K)
|
||||
beta_raw = self.beta_proj(x).view(B, T, HV)
|
||||
o, _ = chunk_kda(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
g_raw,
|
||||
beta_raw,
|
||||
A_log=self.A_log,
|
||||
dt_bias=self.dt_bias,
|
||||
use_qk_l2norm_in_kernel=self.use_qk_l2norm_in_kernel,
|
||||
use_gate_in_kernel=self.use_gate_in_kernel,
|
||||
use_beta_sigmoid_in_kernel=self.use_beta_sigmoid_in_kernel,
|
||||
safe_gate=self.lower_bound is not None,
|
||||
lower_bound=self.lower_bound,
|
||||
chunk_size=self.chunk_size,
|
||||
backend=self.kda_backend,
|
||||
)
|
||||
return self.o_proj(o.reshape(B, T, HV * V))
|
||||
@@ -0,0 +1,118 @@
|
||||
"""Stable LatentMoE (K3): shared 全宽 + routed 半宽专家 + SiTU-GLU + Top-k.
|
||||
|
||||
对照 learning/kimi-k3-notes §Stable LatentMoE:
|
||||
z = W_down(x) [B, T, ℓ] ℓ = d/2 latent 接口宽
|
||||
u = Σ_{i∈Top-k(x)} p_i E_i^rt(z) [B, T, ℓ] routed 专家只在 ℓ 上算
|
||||
y = Σ_j E_j^sh(x) + W_up RMSNorm(u) [B, T, d] shared 全宽
|
||||
|
||||
SiTU-GLU: gate = β1·tanh(W_g x/β1)⊙σ(W_g x); up = β2·tanh(W_u x/β2)
|
||||
||SiTU-GLU||_∞ ≤ β1·β2 (=100), 原点附近≈SwiGLU, 远端软饱和防低精度溢出.
|
||||
E: R^in → R^in (内部中间维 d_ff).
|
||||
|
||||
Router: Top-k logits 基于全宽 x (笔记 Topk(x)); 归一化权重取 softmax(topk).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
from .rmsnorm import RMSNorm
|
||||
|
||||
|
||||
class SiTU(nn.Module):
|
||||
"""SiTU-GLU expert: gate 支软上限 β1, up 支软上限 β2, 输出回到输入维."""
|
||||
|
||||
def __init__(self, dim_in: int, dim_ff: int, beta1: float = 4.0, beta2: float = 25.0):
|
||||
super().__init__()
|
||||
self.beta1, self.beta2 = beta1, beta2
|
||||
self.w_g = nn.Linear(dim_in, dim_ff, bias=False)
|
||||
self.w_u = nn.Linear(dim_in, dim_ff, bias=False)
|
||||
self.w_o = nn.Linear(dim_ff, dim_in, bias=False)
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
wg = self.w_g(x)
|
||||
g = self.beta1 * torch.tanh(wg / self.beta1) * torch.sigmoid(wg)
|
||||
u = self.beta2 * torch.tanh(self.w_u(x) / self.beta2)
|
||||
return self.w_o(g * u)
|
||||
|
||||
|
||||
class LatentMoE(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
latent_size: int,
|
||||
n_routed: int,
|
||||
top_k: int,
|
||||
n_shared: int,
|
||||
d_ff: int,
|
||||
beta1: float = 4.0,
|
||||
beta2: float = 25.0,
|
||||
):
|
||||
super().__init__()
|
||||
self.latent_size = latent_size
|
||||
self.n_routed = n_routed
|
||||
self.top_k = top_k
|
||||
|
||||
self.down = nn.Linear(hidden_size, latent_size, bias=False) # W↓
|
||||
self.router = nn.Linear(hidden_size, n_routed, bias=False) # Top-k logits
|
||||
self.shared = nn.ModuleList(
|
||||
[SiTU(hidden_size, d_ff, beta1, beta2) for _ in range(n_shared)]
|
||||
)
|
||||
self.experts = nn.ModuleList(
|
||||
[SiTU(latent_size, d_ff, beta1, beta2) for _ in range(n_routed)]
|
||||
)
|
||||
self.norm = RMSNorm(latent_size)
|
||||
self.up = nn.Linear(latent_size, hidden_size, bias=False) # W↑
|
||||
self.last_route_ids: torch.Tensor | None = None
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config) -> LatentMoE:
|
||||
return cls(
|
||||
config.hidden_size,
|
||||
config.moe_latent_size,
|
||||
config.n_routed,
|
||||
config.top_k,
|
||||
config.n_shared,
|
||||
config.moe_d_ff,
|
||||
config.situ_beta1,
|
||||
config.situ_beta2,
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
B, T, _ = x.shape
|
||||
z = self.down(x) # [B, T, ℓ]
|
||||
|
||||
logits = self.router(x) # [B, T, n_routed]
|
||||
topk = torch.topk(logits, self.top_k, dim=-1)
|
||||
ids = topk.indices # [B, T, k]
|
||||
self.last_route_ids = ids.detach()
|
||||
probs = F.softmax(topk.values, dim=-1) # [B, T, k]
|
||||
|
||||
# 向量化 routed: 预计算全部专家输出, 按 token 的 Top-k id 取
|
||||
all_out = torch.stack([e(z) for e in self.experts]) # [R, B, T, ℓ]
|
||||
all_out = all_out.permute(1, 2, 0, 3).reshape(B * T, self.n_routed, self.latent_size)
|
||||
u = torch.zeros(B, T, self.latent_size, device=x.device, dtype=x.dtype)
|
||||
for i in range(self.top_k):
|
||||
idx = ids[:, :, i].reshape(B * T) # [B*T]
|
||||
sel = all_out[torch.arange(B * T, device=x.device), idx] # [B*T, ℓ]
|
||||
u += probs[:, :, i : i + 1] * sel.reshape(B, T, self.latent_size)
|
||||
|
||||
shared_out = torch.stack([e(x) for e in self.shared]).sum(0) # [B, T, d]
|
||||
return shared_out + self.up(self.norm(u))
|
||||
|
||||
|
||||
def moe_route_frac(model: nn.Module) -> torch.Tensor | None:
|
||||
"""Mean expert occupancy over LatentMoE layers from the last forward."""
|
||||
hists: list[torch.Tensor] = []
|
||||
n_routed: int | None = None
|
||||
for module in model.modules():
|
||||
if not isinstance(module, LatentMoE) or module.last_route_ids is None:
|
||||
continue
|
||||
n_routed = module.n_routed
|
||||
ids = module.last_route_ids.reshape(-1)
|
||||
hists.append(torch.bincount(ids, minlength=n_routed).float())
|
||||
if not hists or n_routed is None:
|
||||
return None
|
||||
stacked = torch.stack(hists).sum(0)
|
||||
return stacked / stacked.sum().clamp_min(1.0)
|
||||
@@ -0,0 +1,97 @@
|
||||
"""Gated MLA (K3): NoPE, latent KV compression, matrix absorption, full-rank output gate.
|
||||
|
||||
K3 相对 DeepSeek MLA 的三个改动 (对照 learning/kimi-k3-notes):
|
||||
1. NoPE — 不显式 RoPE; 位置感交给夹层 KDA 的 decay/gate。
|
||||
2. 矩阵吸收 — 训练/推理都不解压 K/V: q 吸收 W_UK 后直接与 latent c 内积,
|
||||
输出先在 latent 加权再乘 W_UV 还原 (v2 吸收版)。
|
||||
3. Full-rank 输出门 — y = W_o[ σ(W_g x) ⊙ õ ]。
|
||||
|
||||
形状 (小规模 toy, d 为 hidden):
|
||||
c = RMSNorm(kv_down(x)) [B, T, r] latent
|
||||
q = q_up(RMSNorm(q_down(x))) [B, T, H, d_q] d_q = d_nope (NoPE)
|
||||
W_UK = kv_up[.., :H*d_q].view(H,d_q,r) W_UV = kv_up[.., H*d_q:].view(H,d_v,r)
|
||||
score = (q @ W_UK^T) @ c^T [B, H, T, T] causal
|
||||
õ = (softmax(score) @ c) @ W_UV^T [B, T, H, d_v]
|
||||
y = o_proj( σ(W_g x) ⊙ õ_head ) [B, T, d]
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
from .rmsnorm import RMSNorm
|
||||
|
||||
|
||||
class GatedMLA(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
num_heads: int,
|
||||
kv_lora_rank: int,
|
||||
q_lora_rank: int,
|
||||
qk_nope_head_dim: int,
|
||||
v_head_dim: int,
|
||||
):
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
self.num_heads = num_heads
|
||||
self.qk_nope_head_dim = qk_nope_head_dim
|
||||
self.v_head_dim = v_head_dim
|
||||
|
||||
# Q 低秩路径 (NoPE, 只有 nope 段)
|
||||
self.q_down = nn.Linear(hidden_size, q_lora_rank, bias=False)
|
||||
self.q_norm = RMSNorm(q_lora_rank)
|
||||
self.q_up = nn.Linear(q_lora_rank, num_heads * qk_nope_head_dim, bias=False)
|
||||
|
||||
# KV latent 压缩 + 解压 (W_UK | W_UV 拼接在同一矩阵里)
|
||||
self.kv_down = nn.Linear(hidden_size, kv_lora_rank, bias=False)
|
||||
self.kv_norm = RMSNorm(kv_lora_rank)
|
||||
self.kv_up = nn.Linear(
|
||||
kv_lora_rank, num_heads * (qk_nope_head_dim + v_head_dim), bias=False
|
||||
)
|
||||
|
||||
# Full-rank 输出门: σ(W_g x) 与 õ (H*d_v 维) 逐元素相乘
|
||||
self.gate = nn.Linear(hidden_size, num_heads * v_head_dim, bias=False)
|
||||
self.o_proj = nn.Linear(num_heads * v_head_dim, hidden_size, bias=False)
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config) -> GatedMLA:
|
||||
return cls(
|
||||
config.hidden_size,
|
||||
config.num_heads,
|
||||
config.kv_lora_rank,
|
||||
config.q_lora_rank,
|
||||
config.qk_nope_head_dim,
|
||||
config.v_head_dim,
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
B, T, _ = x.shape
|
||||
H, r = self.num_heads, self.kv_up.in_features
|
||||
|
||||
c = self.kv_norm(self.kv_down(x)) # [B, T, r]
|
||||
q = self.q_up(self.q_norm(self.q_down(x))) # [B, T, H*d_q]
|
||||
q = q.view(B, T, H, self.qk_nope_head_dim) # [B, T, H, d_q]
|
||||
|
||||
w = self.kv_up.weight # [H*(d_q+d_v), r]
|
||||
w_uk = w[: H * self.qk_nope_head_dim].view(H, self.qk_nope_head_dim, r)
|
||||
w_uv = w[H * self.qk_nope_head_dim :].view(H, self.v_head_dim, r)
|
||||
|
||||
# 吸收 W_UK 进 query: score = (q @ W_UK^T) @ c^T
|
||||
q_absorb = torch.einsum("bthd,hdj->bthj", q, w_uk) # [B, T, H, r]
|
||||
scores = torch.einsum("bthj,bsj->bhts", q_absorb, c) # [B, H, T, T]
|
||||
|
||||
mask = torch.triu(
|
||||
torch.ones(T, T, dtype=torch.bool, device=x.device), diagonal=1
|
||||
)
|
||||
scores = scores.masked_fill(mask, float("-inf"))
|
||||
attn = F.softmax(scores, dim=-1) # [B, H, T, T]
|
||||
|
||||
# 先在 latent 加权, 再乘 W_UV^T 还原 v —— 永不解压
|
||||
latent_out = torch.einsum("bhts,bsj->bhtj", attn, c) # [B, H, T, r]
|
||||
o_heads = torch.einsum("bhtj,hvj->bhtv", latent_out, w_uv) # [B, H, T, d_v]
|
||||
|
||||
o_heads = o_heads.transpose(1, 2).reshape(B, T, H * self.v_head_dim)
|
||||
gate = torch.sigmoid(self.gate(x)) # [B, T, H*d_v]
|
||||
return self.o_proj(gate * o_heads) # [B, T, d]
|
||||
@@ -0,0 +1,17 @@
|
||||
"""RMSNorm used by attention, FFN, and the final LM stem."""
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
def __init__(self, dim: int, eps: float = 1e-6):
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
self.eps = eps
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
dtype = x.dtype
|
||||
x = x.float()
|
||||
return (x * torch.rsqrt(x.square().mean(-1, keepdim=True) + self.eps)).to(dtype) * self.weight
|
||||
@@ -0,0 +1,20 @@
|
||||
"""SwiGLU FFN: x [B,T,D] -> y [B,T,D]."""
|
||||
from __future__ import annotations
|
||||
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
|
||||
class SwiGLUMLP(nn.Module):
|
||||
def __init__(self, hidden_size: int, intermediate_size: int):
|
||||
super().__init__()
|
||||
self.w1 = nn.Linear(hidden_size, intermediate_size, bias=False)
|
||||
self.w3 = nn.Linear(hidden_size, intermediate_size, bias=False)
|
||||
self.w2 = nn.Linear(intermediate_size, hidden_size, bias=False)
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config) -> SwiGLUMLP:
|
||||
return cls(config.hidden_size, config.intermediate_size)
|
||||
|
||||
def forward(self, x):
|
||||
return self.w2(F.silu(self.w1(x)) * self.w3(x))
|
||||
Reference in New Issue
Block a user