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
+111
View File
@@ -0,0 +1,111 @@
% teach:
% gap: 读者知道 g_t 是 gate 但不知道它怎么从 raw projection 变成一个负的 log-space 衰减
% takeaway: safe gate 用 sigmoid 把值夹在 [lower_bound, 0], standard gate 用 -softplus 保证负
% jump: 论文没解释为什么需要 A_log 和 dt_bias 两层
% omit: Triton gate kernel 的 fused 实现细节
\section{Gate 激活}
\splabel{C2}
\subsection{Gate 的角色}
回顾 §1:$S_{\mathrm{dec}} = \exp(g_t) \odot S_{t-1}$。$g_t$ 必须 $\leq 0$
才是衰减($\exp(g_t) \leq 1$),否则状态会指数增长爆炸。
\texttt{g\_raw} 是从 \texttt{g\_proj(x)} 出来的 raw 值,没有约束。
Gate 激活函数的任务是:把 raw 值映射到一个保证 $\leq 0$ 的范围。
\subsection{两种 Gate}
\begin{center}
\begin{tabular}{p{3cm}p{5.5cm}p{5cm}}
\toprule
& \textbf{Standard gate} & \textbf{Safe gate} \\
\midrule
公式 &
$g = -\mathrm{rate} \cdot \mathrm{softplus}(\mathrm{input})$ &
$g = L \cdot \sigma(\mathrm{rate} \cdot \mathrm{input})$ \\
值域 &
$(-\infty, 0]$ &
$[L, 0]$($L$ 是 lower\_bound,如 $-5$) \\
衰减范围 &
$\exp(g) \in (0, 1]$ &
$\exp(g) \in [\exp(L), 1]$ \\
稳定性 &
衰减可以任意快 &
衰减有下限,不会瞬间清零 \\
\bottomrule
\end{tabular}
\end{center}
\noindent 其中:
\begin{itemize}[nosep]
\item $\mathrm{input} = g_{\mathrm{raw}} + \Delta_b$ \quad($\Delta_b$
是 \texttt{dt\_bias} \shape{HV, K})
\item $\mathrm{rate} = \exp(A_{\log})$ \quad($A_{\log}$ 是
\texttt{A\_log} \shape{HV},head-wise 可学习)
\end{itemize}
\begin{importantbox}{如果你只记一件事}
Safe gate = $L \cdot \sigma(\mathrm{rate} \cdot \mathrm{input})$,
$L=-5$ 时 $\exp(g) \geq \exp(-5) \approx 0.0067$,
状态永远不会被``一次性清零''。
\end{importantbox}
\subsection{代码对照}
\begin{codemathtop}{ops/reference/gate.py — kda\_gate\_reference}
\begin{lstlisting}
def kda_gate_reference(g, A_log, dt_bias=None, *,
safe_gate=False, lower_bound=None):
HV, K = g.shape[-2:]
gate_input = g if dt_bias is None else g + dt_bias.view(HV, K)
rate = A_log.view(HV, 1).exp()
if safe_gate:
# safe: g in [lower_bound, 0]
return lower_bound * torch.sigmoid(rate * gate_input)
# standard: g in (-inf, 0]
return -rate * F.softplus(gate_input)
\end{lstlisting}
\end{codemathtop}
\subsection{初始化与默认值}
\begin{center}
\begin{tabular}{llp{7cm}}
\toprule
参数 & 初始值 & 效果 \\
\midrule
\texttt{A\_log} & $\mathbf{0}$ \shape{HV} & $\mathrm{rate} = \exp(0) = 1$,不缩放 \\
\texttt{dt\_bias} & $-4.0$ \shape{HV, K} & 初始时 $\mathrm{input} \approx g_{\mathrm{raw}} - 4$,
配合 safe gate ($L=-5$) 得到 $g \approx -5 \cdot \sigma(-4) \approx -0.09$,
即 $\exp(g) \approx 0.91$(约 91\% 状态保留) \\
\texttt{lower\_bound} & $-5.0$ & safe gate 的下限 \\
\bottomrule
\end{tabular}
\end{center}
\subsection{Gate 在 API 中的位置}
Gate 激活在 \texttt{ops/api.py} 的 \texttt{chunk\_kda} 中调用,
在进入 chunkwise 或 recurrent 核心之前完成。
当 \texttt{use\_gate\_in\_kernel=True} 时,\texttt{g\_raw} 进入 API,
API 内部完成 gate 激活;否则调用方自己完成。
\begin{lstlisting}
# ops/api.py (simplified)
if use_gate_in_kernel:
gate_input = g + dt_bias.view(g.shape[-2:])
rate = A_log.exp().view(1, 1, -1, 1)
if safe_gate:
g = lower_bound * torch.sigmoid(rate * gate_input)
else:
g = -rate * F.softplus(gate_input)
\end{lstlisting}
\subsection{本章小结}
Gate 把 raw projection 映射到 $\leq 0$ 的 log-space 衰减率。
Safe gate 用 sigmoid 限制在 $[L, 0]$,防止瞬间清零;
standard gate 用 softplus 不限制下限。默认配置下初始状态保留约 91\%。