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,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 内部完成。
|
||||
Reference in New Issue
Block a user