% 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\%。