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,162 @@
|
||||
% teach:
|
||||
% gap: 读者已知各组件但不知道怎么组装成完整模型
|
||||
% takeaway: K3 = Hybrid(3 KDA + 1 MLA) × DecoderBlock(attn + MoE), 末层强制 MLA
|
||||
% jump: 为什么每 4 层才放一次 MLA?位置感知只需要周期性提供
|
||||
% omit: 0.5b preset 的训练超参
|
||||
|
||||
\section{K3 混合架构}
|
||||
|
||||
\subsection{整体结构}
|
||||
|
||||
\begin{center}
|
||||
\texttt{Embedding} $\to$ \texttt{DecoderBlock} $\times L$ $\to$ \texttt{RMSNorm} $\to$ \texttt{LM Head}
|
||||
\end{center}
|
||||
|
||||
每个 \texttt{DecoderBlock} 是 Pre-Norm 残差:
|
||||
|
||||
\begin{lstlisting}
|
||||
def forward(self, x):
|
||||
x = x + self.attn(self.attn_norm(x)) # mixing
|
||||
return x + self.ffn(self.ffn_norm(x)) # channel
|
||||
\end{lstlisting}
|
||||
|
||||
\subsection{Hybrid Attention Pattern}
|
||||
|
||||
K3 用两种 attention 层交替:
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{ccccccccc}
|
||||
\toprule
|
||||
层 & 0 & 1 & 2 & 3 & 4 & 5 & 6 & 7 \\
|
||||
\midrule
|
||||
Attn & KDA & KDA & KDA & \textbf{MLA} & KDA & KDA & KDA & \textbf{MLA} \\
|
||||
FFN & MoE & MoE & MoE & MoE & MoE & MoE & MoE & MoE \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\noindent 规则:每 4 层放 1 次 MLA(0-based 层 3, 7, 11, ...),\textbf{末层强制 MLA}。
|
||||
|
||||
\begin{codemathtop}{models/k3\_config.py — layer\_types}
|
||||
\begin{lstlisting}
|
||||
def layer_types(self) -> list[str]:
|
||||
"""Hybrid: 3 KDA + 1 MLA per group, last always 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" # last layer forced
|
||||
return types
|
||||
|
||||
def layer_specs(self):
|
||||
return [(kind, "moe") for kind in self.layer_types()]
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\begin{knowledgebox}{为什么 KDA 不需要 RoPE?}
|
||||
KDA 的 gate/decay 机制天然提供位置感知:
|
||||
远的 token 衰减更多,近的保留更多。
|
||||
但 softmax attention(MLA)没有这个机制,所以真实 K3 用 NoPE
|
||||
(本复现的 MLA 也是 NoPE)。
|
||||
位置感知从 KDA 层``渗透''到 MLA 层——3:1 的比例足够了。
|
||||
\end{knowledgebox}
|
||||
|
||||
\subsection{CausalLM 完整数据流}
|
||||
|
||||
\begin{codemathtop}{models/causal\_lm.py — CausalLM}
|
||||
\begin{lstlisting}
|
||||
class CausalLM(nn.Module):
|
||||
def __init__(self, config):
|
||||
self.embedding = nn.Embedding(vocab_size, D) # [V, D]
|
||||
self.blocks = ModuleList([
|
||||
DecoderBlock.from_spec(config, attn, ffn)
|
||||
for attn, ffn in config.layer_specs()
|
||||
])
|
||||
self.norm = RMSNorm(D)
|
||||
self.lm_head = nn.Linear(D, vocab_size) # [V, D]
|
||||
if config.tie_word_embeddings:
|
||||
self.lm_head.weight = self.embedding.weight
|
||||
|
||||
def forward(self, input_ids, labels=None):
|
||||
x = self.embedding(input_ids) # [B,T] -> [B,T,D]
|
||||
if self.mixer is None: # attnres="off"
|
||||
for block in self.blocks:
|
||||
x = block(x) # [B,T,D] -> [B,T,D]
|
||||
else:
|
||||
x = self.mixer(x) # AttnRes 深度残差, 见 §9
|
||||
logits = self.lm_head(self.norm(x)) # [B,T,D] -> [B,T,V]
|
||||
if labels is None:
|
||||
return logits
|
||||
# Shifted CE: predict next token
|
||||
return cross_entropy(logits[:,:-1], labels[:,1:])
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\begin{knowledgebox}{残差流是可替换的}
|
||||
上面的 \texttt{DecoderBlock} 逐层堆叠(\texttt{x = x + sublayer(norm(x))})
|
||||
是 \texttt{config.attnres="off"} 时的默认路径。
|
||||
置为 \texttt{"full"} / \texttt{"block"} 时,\texttt{CausalLM} 会把每个 block
|
||||
拆成 attn / ffn 两个原子子层交给 \texttt{mixer},用\textbf{深度维注意力}
|
||||
代替等权残差加法——见 \S9。真实 K3 用的是 \texttt{block} 模式。
|
||||
\end{knowledgebox}
|
||||
|
||||
\subsection{两种配置}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{lll}
|
||||
\toprule
|
||||
& \textbf{KDAConfig}(纯 KDA) & \textbf{K3Config}(混合) \\
|
||||
\midrule
|
||||
Attn & KDA only & 3 KDA + 1 MLA \\
|
||||
FFN & SwiGLU & LatentMoE \\
|
||||
典型规模 & \textasciitilde8M (toy) & \textasciitilde8M (toy) / \textasciitilde500M (0.5b) \\
|
||||
\texttt{layer\_specs()} & \texttt{[("kda","swiglu")] * L} & \texttt{[(kind,"moe") for kind in ...]} \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{K3 toy 尺寸}
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{llll}
|
||||
\toprule
|
||||
参数 & 真实 K3 & toy 复现 & 缩比 \\
|
||||
\midrule
|
||||
$D$ & 7168 & 256 & 28$\times$ \\
|
||||
$L$ & 93 & 4 & 23$\times$ \\
|
||||
$H = H_V$ & 96 & 8 & 12$\times$ \\
|
||||
$K = V$ & 128 & 16 & 8$\times$ \\
|
||||
kv\_lora\_rank & 512 & 32 & 16$\times$ \\
|
||||
q\_lora\_rank & 1536 & 64 & 24$\times$ \\
|
||||
$\ell$ (MoE latent) & 3584 & 128 & 28$\times$ \\
|
||||
$n_{\mathrm{routed}}$ / Top-$k$ & 896/16 & 16/2 & 56$\times$ / 8$\times$ \\
|
||||
\bottomrule
|
||||
\end{tabular}
|
||||
\end{center}
|
||||
|
||||
\subsection{DecoderBlock 构建}
|
||||
|
||||
\begin{codemathtop}{layers/block.py — build\_attn / build\_ffn}
|
||||
\begin{lstlisting}
|
||||
def build_attn(config, kind: str) -> nn.Module:
|
||||
if kind == "kda": return KDAAttention.from_config(config)
|
||||
if kind == "mla": return GatedMLA.from_config(config)
|
||||
|
||||
def build_ffn(config, kind: str) -> nn.Module:
|
||||
if kind == "swiglu": return SwiGLUMLP.from_config(config)
|
||||
if kind == "moe": return LatentMoE.from_config(config)
|
||||
|
||||
class DecoderBlock(nn.Module):
|
||||
def forward(self, x):
|
||||
x = x + self.attn(self.attn_norm(x))
|
||||
return x + self.ffn(self.ffn_norm(x))
|
||||
\end{lstlisting}
|
||||
\end{codemathtop}
|
||||
|
||||
\subsection{本章小结}
|
||||
|
||||
K3 架构 = Hybrid Attention(3 KDA + 1 MLA,末层强制 MLA)+ LatentMoE。
|
||||
KDA 层提供线性复杂度的序列混合和位置感知(通过 decay),
|
||||
MLA 层提供全局 softmax attention(NoPE,利用 KDA 渗透的位置信息)。
|
||||
每层默认是 Pre-Norm 残差 DecoderBlock;\texttt{config.attnres} 可以把这条
|
||||
等权残差流换成 AttnRes 深度注意力(\S9)。
|
||||
Reference in New Issue
Block a user