Files
K3/notes/sections/sec-02.tex
T
dela 584f7e9e73 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.
2026-08-25 14:43:17 +08:00

112 lines
3.7 KiB
TeX
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
% 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\%。