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
+75
View File
@@ -0,0 +1,75 @@
% teach:
% gap: 读者已知各组件, 但不清楚它们怎么黏在一起成为一个层
% takeaway: KDAAttention = 投影 → gate+norm → chunk_kda → output 投影, 整个层就是 x → y [B,T,D]
% jump: none
% omit: from_config 工厂方法细节
\section{KDAAttention 层}
\subsection{完整数据流}
\texttt{KDAAttention} 把投影、gate 激活、KDA 核心计算和输出投影封装成一个
\texttt{[B,T,D] $\to$ [B,T,D]} 的模块。
\begin{center}
\begin{tabular}{rlll}
\toprule
步骤 & 操作 & 输入形状 & 输出形状 \\
\midrule
1 & \texttt{q\_proj(x)} & \shape{B,T,D} & \shape{B,T,H,K} \\
2 & \texttt{k\_proj(x)} & \shape{B,T,D} & \shape{B,T,H,K} \\
3 & \texttt{v\_proj(x)} & \shape{B,T,D} & \shape{B,T,HV,V} \\
4 & \texttt{g\_proj(x)} & \shape{B,T,D} & \shape{B,T,HV,K} \\
5 & \texttt{beta\_proj(x)} & \shape{B,T,D} & \shape{B,T,HV} \\
6 & \texttt{chunk\_kda(...)} & 上述 5 项 + 参数 & \shape{B,T,HV,V} \\
7 & \texttt{o.reshape(...)} & \shape{B,T,HV,V} & \shape{B,T,HV \cdot V} \\
8 & \texttt{o\_proj(...)} & \shape{B,T,HV \cdot V} & \shape{B,T,D} \\
\bottomrule
\end{tabular}
\end{center}
\subsection{chunk\_kda 内部做了什么}
\texttt{chunk\_kda}(\texttt{ops/api.py})是统一入口,按 \texttt{backend} 分发:
\begin{enumerate}[nosep]
\item 如果 \texttt{use\_qk\_l2norm\_in\_kernel}:$q, k \leftarrow \text{L2-normalize}(q), \text{L2-normalize}(k)$
\item 如果 \texttt{use\_beta\_sigmoid\_in\_kernel}:$\beta \leftarrow \sigma(\beta_{\mathrm{raw}})$
\item 如果 \texttt{use\_gate\_in\_kernel}:应用 gate 激活(§2)
\item 调用 \texttt{naive\_chunk\_kda}(或 triton/fla 版本)
\end{enumerate}
\begin{knowledgebox}{三个 ``in\_kernel'' 开关}
\begin{itemize}[nosep]
\item \texttt{use\_qk\_l2norm}:L2-norm 让 $\langle q, k \rangle$ 变成余弦相似度,
稳定训练
\item \texttt{use\_beta\_sigmoid}:sigmoid 把 $\beta$ 限制在 $(0,1)$,
控制写入强度
\item \texttt{use\_gate\_in\_kernel}:gate 激活在 API 内部完成(vs 调用方自己做)
\end{itemize}
默认三个都是 \texttt{True}。
\end{knowledgebox}
\subsection{可学习参数清单}
\begin{center}
\begin{tabular}{lll}
\toprule
参数 & 形状 & 说明 \\
\midrule
\texttt{q\_proj.weight} & \shape{H \cdot K, D} & query 投影 \\
\texttt{k\_proj.weight} & \shape{H \cdot K, D} & key 投影 \\
\texttt{v\_proj.weight} & \shape{HV \cdot V, D} & value 投影 \\
\texttt{g\_proj.weight} & \shape{HV \cdot K, D} & gate 投影 \\
\texttt{beta\_proj.weight} & \shape{HV, D} & beta 投影 \\
\texttt{o\_proj.weight} & \shape{D, HV \cdot V} & 输出投影 \\
\texttt{A\_log} & \shape{HV} & head-wise 衰减率(log-space)\\
\texttt{dt\_bias} & \shape{HV, K} & per-dim gate bias \\
\bottomrule
\end{tabular}
\end{center}
\subsection{本章小结}
KDAAttention 是一个完整的 mixing 模块:5 个线性投影 + gate 激活 + KDA 核心 + 输出投影。
三个 ``in\_kernel'' 开关控制 L2-norm、sigmoid、gate 是否在 API 内部完成。