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
+127
View File
@@ -0,0 +1,127 @@
% teach:
% gap: 读者知道 softmax attention 但不知道线性注意力怎么维护状态矩阵
% takeaway: KDA 用 delta rule 逐步更新 [K,V] 状态矩阵, 写入=擦旧写新, 每步 O(KV)
% jump: 为什么 r_t = v - k·S 而不是直接用 v?delta rule 的"先擦再写"
% omit: KDA 论文的 related work、实验细节
\section{KDA 递归核心}
\splabel{C1}
\subsection{动机:从 softmax 到状态矩阵}
标准 attention 每个 token 都要回看所有历史,复杂度 $O(T^2)$。
线性注意力换掉 softmax,把 $\sum_j v_j k_j^T$ 压成一个 $K \times V$ 的状态矩阵 $S$,
每步只做 $o_t = q_t \cdot S$,复杂度降到 $O(T \cdot K \cdot V)$。
但裸线性注意力的问题是:$S$ 只能加,不能改。写进去的信息永远在那里。
KDA 的核心想法是给 $S$ 加两个操作:\textbf{衰减}(逐渐忘记旧信息)和
\textbf{delta rule}(先擦旧的,再写新的)。
\begin{importantbox}{如果你只记一件事}
KDA 的状态更新 = 衰减旧状态 + $\beta_t k_t$ 写入 $(v_t - k_t \cdot S_{\mathrm{dec}})$。
减去 $k_t \cdot S_{\mathrm{dec}}$ 就是"先把 $k_t$ 方向的旧预测擦掉"。
\end{importantbox}
\subsection{逐步公式}
\noindent\textbf{输入张量:}
\begin{center}
\begin{tabular}{lll}
\toprule
符号 & 形状 & 含义 \\
\midrule
$q_t$ & \shape{B, HV, K} & query(已经 repeat\_interleave 到 HV) \\
$k_t$ & \shape{B, HV, K} & key(同上) \\
$v_t$ & \shape{B, HV, V} & value \\
$g_t$ & \shape{B, HV, K} & gate(log-space 衰减率,逐维) \\
$\beta_t$ & \shape{B, HV} & 写入强度标量 \\
$S_{t-1}$ & \shape{B, HV, K, V} & 上一步的 KV 状态 \\
\bottomrule
\end{tabular}
\end{center}
\noindent\textbf{四步更新:}
\begin{enumerate}[leftmargin=2em]
\item \textbf{衰减旧状态}(逐元素,$g_t$ 是 log-space 所以取 exp):
\[
S_{\mathrm{dec}} = \exp(g_t) \odot S_{t-1}
\qquad \shape{B, HV, K, V}
\]
\item \textbf{计算残差}(先用 $k_t$ 查旧状态,得到"旧预测",再减掉):
\[
p_t = \sum_k k_{t,k} \cdot S_{\mathrm{dec},k,\cdot}
= \texttt{einsum('bhk, bhkv -> bhv')}
\qquad \shape{B, HV, V}
\]
\[
r_t = v_t - p_t \qquad \shape{B, HV, V}
\]
\item \textbf{写入状态}(外积 rank-1 更新):
\[
a_t = \beta_t \cdot k_t \qquad \shape{B, HV, K}
\]
\[
S_t = S_{\mathrm{dec}} + a_t \otimes r_t
= S_{\mathrm{dec}} + \texttt{einsum('bhk, bhv -> bhkv')}
\qquad \shape{B, HV, K, V}
\]
\item \textbf{读出}:
\[
o_t = \frac{1}{\sqrt{K}} \cdot q_t \cdot S_t
= \texttt{einsum('bhk, bhkv -> bhv')}
\qquad \shape{B, HV, V}
\]
\end{enumerate}
\subsection{代码对照}
\begin{codemathtop}{ops/reference/recurrent.py — naive\_kda\_fwd (核心循环)}
\begin{lstlisting}
for t in range(T):
q_t = qe[:, t] # [B, HV, K]
k_t = ke[:, t] # [B, HV, K]
v_t = v[:, t] # [B, HV, V]
g_t = g[:, t] # [B, HV, K]
b_t = beta[:, t] # [B, HV]
# Step 1: decay
S_dec = S * g_t.exp().unsqueeze(-1) # [B,HV,K,V]
# Step 2: residual (delta rule)
p_t = einsum('bhk, bhkv -> bhv', k_t, S_dec)
r_t = v_t - p_t # [B,HV,V]
# Step 3: write (rank-1 update)
a_t = b_t.unsqueeze(-1) * k_t # [B,HV,K]
S = S_dec + einsum('bhk, bhv -> bhkv', a_t, r_t)
# Step 4: read
o[:, t] = einsum('bhk, bhkv -> bhv', q_t, S)
\end{lstlisting}
\end{codemathtop}
\begin{warningbox}{为什么 exp(g\_t) 要 unsqueeze(-1)?}
$g_t$ 的形状是 \shape{B, HV, K},而 $S$ 是 \shape{B, HV, K, V}。
衰减是在 $K$ 维上逐元素(同一 $k$ 索引的所有 $v$ 维度共享同一个衰减率),
所以 \texttt{exp(g\_t).unsqueeze(-1)} 把 K 维 broadcast 到 $K \times V$。
\end{warningbox}
\subsection{Delta rule 的直觉}
\begin{knowledgebox}{为什么减去 $k_t \cdot S_{\mathrm{dec}}$?}
把 $S$ 想象成一个 $K \to V$ 的线性映射。用 $k_t$ 去查它,得到的 $p_t = k_t^T S$
就是``旧状态对 $k_t$ 方向的预测''。如果 $p_t$ 已经很接近 $v_t$,说明这个方向的信息
已经写好了,不需要再写。$r_t = v_t - p_t$ 就是``需要修正的量''。
这就是 Widrow-Hoff delta rule:不是盲目地加,而是只修正误差。
\end{knowledgebox}
\subsection{本章小结}
KDA 的状态更新 = 衰减旧状态 + $\beta_t k_t$ 写入残差 $(v_t - k_t \cdot S_{\mathrm{dec}})$。
每步复杂度 $O(K \cdot V)$(两次矩阵-向量乘 + 一次外积),不需要 softmax。
+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\%。
+148
View File
@@ -0,0 +1,148 @@
% teach:
% gap: 读者知道递归形式但不知道怎么在 GPU 上并行, 以为只能一步步算
% takeaway: chunk 内用下三角线性系统并行求解, chunk 间递推状态, 数值等价于 naive recurrent
% jump: 论文直接写了 triangular solve 但没解释为什么要 solve 而不是直接矩阵乘
% omit: Triton 实现细节
\section{分块并行计算(Chunkwise)}
\splabel{C3}
\subsection{为什么需要分块?}
Naive recurrent 一步一步算,$T$ 步串行,GPU 利用率低。
分块的想法是把序列切成 $T/C$ 个长度为 $C$ 的 chunk:
\begin{itemize}[nosep]
\item \textbf{chunk 内}:$C$ 个 token 之间的依赖可以用矩阵运算并行处理
\item \textbf{chunk 间}:状态 $S$ 从上一个 chunk 传到下一个,仍然是递推
\end{itemize}
\begin{importantbox}{如果你只记一件事}
Chunkwise = chunk 内并行 + chunk 间递推。数值结果与 naive recurrent 逐位一致。
\end{importantbox}
\subsection{chunk 内的 cumsum 与下三角解}
在每个 chunk 内,先对 $g$ 做 cumsum(前缀和),这样衰减就变成了相对距离的函数:
\[
g_{\mathrm{cum},i} = \sum_{j=0}^{i} g_j, \qquad
\text{token } i \text{ 对 token } j \text{ 的衰减} = \exp(g_{\mathrm{cum},i} - g_{\mathrm{cum},j})
\]
定义 \textbf{decayed dot} 矩阵(chunk 内 $C \times C$):
\[
A_{ij} = \langle x_i, \exp(g_{\mathrm{cum},i} - g_{\mathrm{cum},j}) \cdot k_j \rangle
\qquad \shape{..., C, C}
\]
这个矩阵的构造是 chunk 内计算的核心。用它可以构造一个下三角线性系统:
\[
M = I + \mathrm{tril}(A_{kk} \cdot \beta, \text{diagonal}=-1)
\qquad \shape{..., C, C}
\]
\[
M \cdot W = \exp(g_{\mathrm{cum}}) \cdot k \qquad \Rightarrow \qquad
W = M^{-1} (\exp(g_{\mathrm{cum}}) \cdot k)
\]
\[
M \cdot U = v \qquad \Rightarrow \qquad U = M^{-1} v
\]
\subsection{代码对照}
\begin{codemathtop}{ops/reference/chunkwise.py — naive\_chunk\_kda (核心)}
\begin{lstlisting}
# Rearrange: [B,T,H,K] -> [B,H,N,C,K] where N=T/C
q, k = [rearrange(x, 'b (n c) h d -> b h n c d', c=C)
.repeat_interleave(HV//H, dim=1) for x in (q, k)]
v, g = [rearrange(x, 'b (n c) h d -> b h n c d', c=C)
for x in (v, g)]
beta = rearrange(beta, 'b (n c) h -> b h n c', c=C)
q = q * scale
g = g.cumsum(dim=-2) # chunk 内 cumsum
# Construct triangular system
A_kk = _decayed_dot(k, k, g) # [B,HV,N,C,C]
M = eye + (A_kk * beta[...,None,:]).masked_fill(mask_upper, 0)
W = solve_triangular(M, g.exp() * k, upper=False)
U = solve_triangular(M, v, upper=False)
# A_qk: query 对 key 的 decayed dot (含对角线)
A_qk = (_decayed_dot(q, k, g) * beta[...,None,:])
.masked_fill(mask_strict_upper, 0)
\end{lstlisting}
\end{codemathtop}
\subsection{chunk 间递推}
每个 chunk 内算完后,用 $W$ 和 $U$ 来处理跨 chunk 的状态:
\begin{codemathtop}{ops/reference/chunkwise.py — chunk 间循环}
\begin{lstlisting}
S = zeros(B, HV, K, V) # inter-chunk state
for n in range(T // C):
# r = "local residual, adjusted by cross-chunk state"
r = U[:,:,n] - W[:,:,n] @ S # [B,HV,C,V]
# output: cross-chunk part + intra-chunk part
o[:,:,n] = (q_n * g_n.exp()) @ S + A_qk[:,:,n] @ r
# update cross-chunk state
decay = (g_n[:,:,-1:,:] - g_n).exp() # decay to chunk end
S = S * g_n[:,:,-1,:,None].exp() # decay old state
S = S + (decay * k_n).T @ (r * beta_n) # write new
\end{lstlisting}
\end{codemathtop}
\subsection{形状流水线}
\begin{center}
\begin{tabular}{lll}
\toprule
变量 & 形状 & 说明 \\
\midrule
\texttt{q, k} (chunked) & \shape{B, HV, N, C, K} & $N = T/C$ 个 chunk \\
\texttt{v, g} (chunked) & \shape{B, HV, N, C, V/K} & \\
\texttt{beta} (chunked) & \shape{B, HV, N, C} & \\
\texttt{A\_kk} & \shape{B, HV, N, C, C} & key-key decayed dot \\
\texttt{M} & \shape{B, HV, N, C, C} & 下三角系统 \\
\texttt{W} & \shape{B, HV, N, C, K} & $M^{-1}(\exp(g) \cdot k)$ \\
\texttt{U} & \shape{B, HV, N, C, V} & $M^{-1} v$ \\
\texttt{A\_qk} & \shape{B, HV, N, C, C} & query-key decayed dot \\
\texttt{S} & \shape{B, HV, K, V} & 跨 chunk 状态 \\
\texttt{r} & \shape{B, HV, C, V} & 调整后的残差 \\
\texttt{o (chunk n)} & \shape{B, HV, C, V} & 本 chunk 输出 \\
\bottomrule
\end{tabular}
\end{center}
\begin{warningbox}{为什么用 triangular solve 而不是直接矩阵乘?}
Delta rule 的"先擦再写"引入了 chunk 内 token 之间的递归依赖:
token $i$ 的写入依赖 token $j < i$ 的写入结果。
这个依赖关系恰好形成一个下三角线性系统 $M \cdot x = b$,
用 \texttt{solve\_triangular} 可以在 $O(C^2)$ 内并行求解,
而展开递归需要 $O(C)$ 步串行。
\end{warningbox}
\subsection{Decayed dot 函数}
\begin{codemathtop}{ops/reference/chunkwise.py — \_decayed\_dot}
\begin{lstlisting}
def _decayed_dot(x, k, g):
"""A[..., i, j] = <x_i, exp(g_i - g_j) * k_j>"""
C = x.shape[-2]
A = empty(*x.shape[:-2], C, C)
for i in range(C):
decay = (g[..., i:i+1, :] - g).exp() # [.., 1, K] - [.., C, K]
A[..., i, :] = einsum('...jk,...jk->...j',
x[..., i, None, :] * decay, k)
return A
\end{lstlisting}
\end{codemathtop}
\noindent 这是一个 $C \times C$ 的矩阵,每个元素 $(i,j)$ 是
$x_i$ 和 $\exp(g_i - g_j) \cdot k_j$ 的内积。Triton 实现会把这个双循环融合成一个 kernel。
\subsection{本章小结}
分块把 $T$ 步串行拆成 $T/C$ 个 chunk,chunk 内用下三角 solve 并行处理 delta rule 依赖,
chunk 间递推状态 $S$。最终输出与 naive recurrent 逐位相同。
+114
View File
@@ -0,0 +1,114 @@
% teach:
% gap: 读者不知道 q/k 和 v 为什么可以有不同的头数, 以及 repeat_interleave 的反传怎么做
% takeaway: GVA 让 G 组 value heads 共享一组 q/k, forward repeat_interleave, backward sum
% jump: 论文没解释为什么反传是 sum 而不是 mean
% omit: GQA 的历史
\section{GVA(分组值注意力)}
\splabel{GVA}
\subsection{为什么头数不一样?}
标准 MHA 里 $H_q = H_k = H_v$。GQA(Grouped Query Attention)让多组 q/k 共享同一组 v/k,
减少 KV cache。KDA 反过来做:$H$ 组 q/k 对应 $H_V = G \cdot H$ 组 value heads。
直觉:value 维度决定表达能力,多一点 value head 增加容量;
q/k 主要负责路由(``看哪里''),可以共享。
\begin{center}
\begin{tabular}{lll}
\toprule
& 标准头数 & GVA \\
\midrule
$q, k$ & \shape{B, T, H, K} & \shape{B, T, H, K}(不变)\\
$v$ & \shape{B, T, H, V} & \shape{B, T, HV, V}($H_V = G \cdot H$) \\
$g, \beta$ & \shape{B, T, H, K/1} & \shape{B, T, HV, K/1} \\
$S$ & \shape{B, H, K, V} & \shape{B, HV, K, V} \\
\bottomrule
\end{tabular}
\end{center}
\subsection{Forward: repeat\_interleave}
进入 KDA 核心前,$q$ 和 $k$ 从 $H$ 维复制到 $H_V$ 维:
\begin{lstlisting}
G = HV // H
qe = q.repeat_interleave(G, dim=2) * scale # [B,T,H,K] -> [B,T,HV,K]
ke = k.repeat_interleave(G, dim=2) # [B,T,H,K] -> [B,T,HV,K]
\end{lstlisting}
\noindent 例如 $H=4, G=2, H_V=8$:head 0 的 q/k 复制到 value head 0 和 1,
head 1 复制到 value head 2 和 3,依此类推。
\subsection{Backward: view + sum}
反传时,$dq_e$ 和 $dk_e$ 的形状是 \shape{B, T, HV, K}(在 $H_V$ 维上计算的梯度)。
因为 forward 是复制,反传就是求和:
\begin{lstlisting}
# Backward: HV -> H
dq_H = dq_e.view(B, T, H, G, K).sum(dim=3) # [B,T,HV,K] -> [B,T,H,K]
dk_H = dk_e.view(B, T, H, G, K).sum(dim=3)
\end{lstlisting}
\begin{warningbox}{为什么是 sum 不是 mean?}
\texttt{repeat\_interleave} 是\textbf{复制}:$y_0 = x_0, y_1 = x_0, y_2 = x_1, \ldots$
对 $x_0$ 的梯度 = $\frac{\partial L}{\partial y_0} + \frac{\partial L}{\partial y_1}$
= \textbf{sum}(不是 mean)。
这和 \texttt{.expand()} 的反传一样:复制的反传是求和。
\end{warningbox}
\subsection{scale 的处理}
$q$ 在 repeat\_interleave 之后乘了 \texttt{scale = $1/\sqrt{K}$}。
反传时 chain rule 要求 $dq_{\mathrm{orig}} = dq_e \cdot \texttt{scale}$:
\begin{lstlisting}
# q 在 forward 内被乘过 scale, chain rule:
dq_H = dq_H * scale
\end{lstlisting}
\subsection{KDAAttention 层中的投影}
\begin{codemathtop}{layers/kda\_attn.py — forward}
\begin{lstlisting}
def forward(self, x): # x: [B, T, D]
B, T, _ = x.shape
H, HV, K, V = self.num_heads, self.num_value_heads, ...
q = self.q_proj(x).view(B, T, H, K) # [B,T,D] -> [B,T,H*K] -> [B,T,H,K]
k = self.k_proj(x).view(B, T, H, K) # 同上
v = self.v_proj(x).view(B, T, HV, V) # [B,T,D] -> [B,T,HV*V] -> [B,T,HV,V]
g_raw = self.g_proj(x).view(B, T, HV, K)
beta_raw = self.beta_proj(x).view(B, T, HV)
o, _ = chunk_kda(q, k, v, g_raw, beta_raw, ...)
return self.o_proj(o.reshape(B, T, HV * V)) # [B,T,HV,V] -> [B,T,D]
\end{lstlisting}
\end{codemathtop}
\subsection{投影矩阵形状总览}
\begin{center}
\begin{tabular}{llll}
\toprule
投影 & 权重形状 & 输入 & 输出 \\
\midrule
\texttt{q\_proj} & \shape{H \cdot K, D} & \shape{B,T,D} & \shape{B,T,H,K} \\
\texttt{k\_proj} & \shape{H \cdot K, D} & \shape{B,T,D} & \shape{B,T,H,K} \\
\texttt{v\_proj} & \shape{HV \cdot V, D} & \shape{B,T,D} & \shape{B,T,HV,V} \\
\texttt{g\_proj} & \shape{HV \cdot K, D} & \shape{B,T,D} & \shape{B,T,HV,K} \\
\texttt{beta\_proj} & \shape{HV, D} & \shape{B,T,D} & \shape{B,T,HV} \\
\texttt{o\_proj} & \shape{D, HV \cdot V} & \shape{B,T,HV \cdot V} & \shape{B,T,D} \\
\bottomrule
\end{tabular}
\end{center}
\subsection{本章小结}
GVA 让 $H_V = G \cdot H$ 组 value heads 共享 $H$ 组 q/k。
Forward 用 \texttt{repeat\_interleave} 复制,backward 用 \texttt{view+sum} 归约。
$v, g, \beta$ 直接在 $H_V$ 维投影,q/k 在 $H$ 维投影。
+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 内部完成。
+173
View File
@@ -0,0 +1,173 @@
% teach:
% gap: 读者知道标准 MHA 但不知道 MLA 怎么压缩 KV、矩阵吸收怎么避免解压
% takeaway: MLA 把 KV 压成低秩 latent c, 通过吸收 W_UK 进 q 直接在 latent 空间算 attention
% jump: 为什么可以先在 latent 加权再乘 W_UV?因为矩阵乘和加权求和可交换
% omit: RoPE (K3 用 NoPE)
\section{Gated MLA(矩阵吸收版)}
\splabel{C4}
\subsection{标准 MHA 的 KV cache 问题}
标准 MHA 推理时需要缓存所有历史 token 的 $K, V$,cache 大小 $\propto T \cdot H \cdot d$。
MLA 的想法:把 $K, V$ 压缩成一个低秩 latent $c$,cache 大小 $\propto T \cdot r$,
其中 $r \ll H \cdot d$。
\subsection{低秩压缩}
\[
c = \mathrm{RMSNorm}(W_{\downarrow} \cdot x) \qquad \shape{B, T, r}
\]
推理时只缓存 $c$,不缓存解压后的 $K, V$。
解压矩阵 $W_{\mathrm{KV}\uparrow}$ 包含两部分:
\[
W_{\mathrm{KV}\uparrow} = \begin{bmatrix} W_{UK} \\ W_{UV} \end{bmatrix}
\qquad \shape{H \cdot (d_q + d_v), r}
\]
拆开:$W_{UK} \in \mathbb{R}^{H \times d_q \times r}$(key 解压),
$W_{UV} \in \mathbb{R}^{H \times d_v \times r}$(value 解压)。
\subsection{矩阵吸收的核心思路}
\textbf{不解压} $K$ 和 $V$。标准做法会先解压再算 attention:
\begin{center}
\textit{标准}:$k_h = c \cdot W_{UK,h}^T$ \shape{B,T,d_q},
$\mathrm{score} = q_h \cdot k_h^T$
\end{center}
矩阵吸收反过来:把 $W_{UK}$ 吸收进 $q$:
\begin{center}
\textit{吸收}:$q_{\mathrm{abs},h} = q_h \cdot W_{UK,h}$ \shape{B,T,r},
$\mathrm{score} = q_{\mathrm{abs},h} \cdot c^T$
\end{center}
\begin{importantbox}{如果你只记一件事}
$(q \cdot W_{UK}^T) \cdot c^T = q \cdot (W_{UK}^T \cdot c^T) = q_{\mathrm{abs}} \cdot c^T$
吸收后,attention 直接在 latent 空间 $r$ 维上算,永不解压到 $H \cdot d_q$ 维。
\end{importantbox}
\subsection{完整计算流(四步)}
\begin{enumerate}[leftmargin=2em]
\item \textbf{Q 低秩路径}(NoPE,只有 nope 段):
\[
q = W_{q\uparrow} \cdot \mathrm{RMSNorm}(W_{q\downarrow} \cdot x)
\qquad \shape{B, T, H, d_q}
\]
\item \textbf{吸收 $W_{UK}$ + 打分}:
\[
q_{\mathrm{abs}} = q \cdot W_{UK} \quad
\xrightarrow{\texttt{einsum('bthd,hdj->bthj')}} \quad \shape{B, T, H, r}
\]
\[
\mathrm{score} = q_{\mathrm{abs}} \cdot c^T \quad
\xrightarrow{\texttt{einsum('bthj,bsj->bhts')}} \quad \shape{B, H, T, T}
\]
\[
\mathrm{attn} = \mathrm{softmax}(\mathrm{causal\_mask}(\mathrm{score}))
\qquad \shape{B, H, T, T}
\]
\item \textbf{先在 latent 加权,再乘 $W_{UV}^T$}:
\[
\tilde{o}_{\mathrm{lat}} = \mathrm{attn} \cdot c \quad
\xrightarrow{\texttt{einsum('bhts,bsj->bhtj')}} \quad \shape{B, H, T, r}
\]
\[
\tilde{o} = \tilde{o}_{\mathrm{lat}} \cdot W_{UV}^T \quad
\xrightarrow{\texttt{einsum('bhtj,hvj->bhtv')}} \quad \shape{B, H, T, d_v}
\]
\item \textbf{输出门 + 投影}:
\[
y = W_o \big[ \sigma(W_g \cdot x) \odot \tilde{o}_{\mathrm{flat}} \big]
\qquad \shape{B, T, D}
\]
\end{enumerate}
\subsection{代码对照}
\begin{codemathtop}{layers/mla.py — GatedMLA.forward}
\begin{lstlisting}
def forward(self, x): # x: [B, T, D]
B, T, _ = x.shape
H, r = self.num_heads, self.kv_up.in_features
# Step 1: latent + query
c = self.kv_norm(self.kv_down(x)) # [B, T, r]
q = self.q_up(self.q_norm(self.q_down(x))) # [B, T, H*d_q]
q = q.view(B, T, H, self.qk_nope_head_dim) # [B, T, H, d_q]
# Split W_UK, W_UV from kv_up.weight
w = self.kv_up.weight # [H*(d_q+d_v), r]
w_uk = w[:H*d_q].view(H, d_q, r) # [H, d_q, r]
w_uv = w[H*d_q:].view(H, d_v, r) # [H, d_v, r]
# Step 2: absorb W_UK, score
q_absorb = einsum('bthd,hdj->bthj', q, w_uk) # [B,T,H,r]
scores = einsum('bthj,bsj->bhts', q_absorb, c) # [B,H,T,T]
scores = scores.masked_fill(causal_mask, -inf)
attn = softmax(scores, dim=-1) # [B,H,T,T]
# Step 3: latent-space weighted sum, then W_UV
latent_out = einsum('bhts,bsj->bhtj', attn, c) # [B,H,T,r]
o_heads = einsum('bhtj,hvj->bhtv', latent_out, w_uv) # [B,H,T,d_v]
# Step 4: output gate
o_heads = o_heads.transpose(1,2).reshape(B,T, H*d_v)
gate = sigmoid(self.gate(x)) # [B,T,H*d_v]
return self.o_proj(gate * o_heads) # [B,T,D]
\end{lstlisting}
\end{codemathtop}
\begin{warningbox}{为什么可以先加权再乘 $W_{UV}$?}
标准做法:$o = \mathrm{attn} \cdot V = \mathrm{attn} \cdot (c \cdot W_{UV}^T)$
交换顺序:$o = (\mathrm{attn} \cdot c) \cdot W_{UV}^T$
这能成立是因为矩阵乘法的结合律:$A(BC) = (AB)C$。
$\mathrm{attn} \cdot c$ 先在 latent 空间 $r$ 维上加权求和,
得到的 \shape{B,H,T,r} 再乘 $W_{UV}^T$ 还原到 $d_v$ 维。
全程不需要显式构造 $H \cdot T$ 大小的 $V$ 矩阵。
\end{warningbox}
\subsection{形状与参数对比}
\begin{center}
\begin{tabular}{lll}
\toprule
参数 & 形状 & 说明 \\
\midrule
\texttt{kv\_down.weight} & \shape{r, D} & KV latent 压缩 \\
\texttt{kv\_up.weight} & \shape{H \cdot (d_q+d_v), r} & 包含 $W_{UK}$ 和 $W_{UV}$ \\
\texttt{q\_down.weight} & \shape{r_q, D} & Q 低秩 \\
\texttt{q\_up.weight} & \shape{H \cdot d_q, r_q} & Q 解压 \\
\texttt{gate.weight} & \shape{H \cdot d_v, D} & 输出门 \\
\texttt{o\_proj.weight} & \shape{D, H \cdot d_v} & 输出投影 \\
\bottomrule
\end{tabular}
\end{center}
\noindent KV cache 大小对比(推理时):
\begin{center}
\begin{tabular}{ll}
\toprule
方法 & Cache 大小 per token \\
\midrule
标准 MHA & $2 \times H \times d = 2 H d$ \\
MLA (latent) & $r$(只存 $c$) \\
\bottomrule
\end{tabular}
\end{center}
\subsection{本章小结}
Gated MLA 把 KV 压缩到低秩 latent $c$ \shape{B,T,r},通过矩阵吸收
($q_{\mathrm{abs}} = q \cdot W_{UK}$)直接在 latent 空间打分和加权,
永不解压 K/V。输出通过 sigmoid 门控。NoPE:不使用 RoPE,位置感交给夹层 KDA 的 decay/gate。
+158
View File
@@ -0,0 +1,158 @@
% teach:
% gap: 读者知道 MoE 的 top-k 路由但不知道 LatentMoE 的 latent 接口和 SiTU-GLU
% takeaway: LatentMoE 通过 latent 接口把 routed 专家限制在 ℓ=d/2 上算, SiTU-GLU 用软上限防溢出
% jump: 为什么 routed 专家在 latent 空间而 shared 在全宽?省参数
% omit: load balancing loss
\section{SiTU-GLU 与 Stable LatentMoE}
\splabel{C5}
\subsection{SiTU-GLU:带软上限的激活}
SwiGLU 在低精度(fp16/bf16)训练时可能溢出:$\mathrm{silu}(x) \cdot x$ 没有上限。
SiTU-GLU 用 $\tanh$ 给门控和上投影加软上限:
\[
\mathrm{SiTU}(x) = W_o \big[\underbrace{\beta_1 \tanh\!\left(\frac{W_g x}{\beta_1}\right) \cdot \sigma(W_g x)}_{\text{gate}} \;\cdot\; \underbrace{\beta_2 \tanh\!\left(\frac{W_u x}{\beta_2}\right)}_{\text{up}}\big]
\]
\begin{center}
\begin{tabular}{lp{8cm}}
\toprule
性质 & 说明 \\
\midrule
输出上限 & $\|\mathrm{SiTU}\|_\infty \leq \beta_1 \cdot \beta_2 = 4 \times 25 = 100$ \\
原点附近 & $\tanh(x/\beta) \approx x/\beta$,所以 $\beta \cdot \tanh(x/\beta) \approx x$,退化为 SwiGLU \\
远端 & 软饱和,防 fp16 溢出 \\
\bottomrule
\end{tabular}
\end{center}
\begin{codemathtop}{layers/latent\_moe.py — SiTU}
\begin{lstlisting}
class SiTU(nn.Module):
def __init__(self, dim_in, dim_ff, beta1=4.0, beta2=25.0):
self.w_g = nn.Linear(dim_in, dim_ff, bias=False)
self.w_u = nn.Linear(dim_in, dim_ff, bias=False)
self.w_o = nn.Linear(dim_ff, dim_in, bias=False)
def forward(self, x): # [*, dim_in]
wg = self.w_g(x)
g = self.beta1 * tanh(wg / self.beta1) * sigmoid(wg) # gate
u = self.beta2 * tanh(self.w_u(x) / self.beta2) # up
return self.w_o(g * u) # [*, dim_in]
\end{lstlisting}
\end{codemathtop}
\subsection{LatentMoE 架构}
\begin{center}
\begin{tabular}{rl}
\toprule
组件 & 说明 \\
\midrule
\textbf{Shared 专家} & $n_{\mathrm{shared}}$ 个 SiTU,全宽 $d \to d$,所有 token 都经过 \\
\textbf{Routed 专家} & $n_{\mathrm{routed}}$ 个 SiTU,半宽 $\ell \to \ell$($\ell = d/2$) \\
\textbf{Latent 接口} & $W_\downarrow: d \to \ell$, $W_\uparrow: \ell \to d$(压缩/还原) \\
\textbf{Router} & $W_r: d \to n_{\mathrm{routed}}$,Top-k 选择 + softmax 归一化 \\
\bottomrule
\end{tabular}
\end{center}
\subsection{计算流(五步)}
\begin{enumerate}[leftmargin=2em]
\item \textbf{Latent 压缩}:
\[
z = W_\downarrow \cdot x \qquad \shape{B, T, \ell}
\]
\item \textbf{Routing}:
\[
\mathrm{logits} = W_r \cdot x \qquad \shape{B, T, n_{\mathrm{routed}}}
\]
\[
\mathrm{ids}, \mathrm{probs} = \mathrm{TopK}(\mathrm{logits}, k)
\qquad \mathrm{ids}: \shape{B, T, k}, \;\; \mathrm{probs}: \shape{B, T, k}
\]
\item \textbf{Routed 专家}(在 latent 空间 $\ell$ 上):
\[
u = \sum_{i \in \mathrm{Top\text{-}k}} p_i \cdot E_i^{\mathrm{rt}}(z)
\qquad \shape{B, T, \ell}
\]
\item \textbf{Shared 专家}(全宽 $d$):
\[
s = \sum_j E_j^{\mathrm{sh}}(x) \qquad \shape{B, T, d}
\]
\item \textbf{合并}:
\[
y = s + W_\uparrow \cdot \mathrm{RMSNorm}(u) \qquad \shape{B, T, d}
\]
\end{enumerate}
\begin{importantbox}{如果你只记一件事}
Routed 专家只在 $\ell = d/2$ 的 latent 空间操作,
参数量是全宽专家的 $1/4$($\ell^2$ vs $d^2$)。
Shared 专家保持全宽 $d$,提供基础表达能力。
\end{importantbox}
\subsection{代码对照}
\begin{codemathtop}{layers/latent\_moe.py — LatentMoE.forward}
\begin{lstlisting}
def forward(self, x): # [B, T, d]
z = self.down(x) # [B, T, ell]
logits = self.router(x) # [B, T, n_routed]
topk = torch.topk(logits, self.top_k, dim=-1)
ids = topk.indices # [B, T, k]
probs = F.softmax(topk.values, dim=-1) # [B, T, k]
# All expert outputs (vectorized)
all_out = stack([e(z) for e in self.experts]) # [R, B, T, ell]
# Gather top-k and weighted sum
u = zeros(B, T, ell)
for i in range(self.top_k):
idx = ids[:,:,i].reshape(B*T)
sel = all_out[arange, idx]
u += probs[:,:,i:i+1] * sel.reshape(B, T, ell)
shared_out = stack([e(x) for e in self.shared]).sum(0) # [B, T, d]
return shared_out + self.up(self.norm(u)) # [B, T, d]
\end{lstlisting}
\end{codemathtop}
\begin{warningbox}{为什么 router 用 $x$(全宽)而不是 $z$(latent)?}
路由需要看到 token 的完整表示才能做好选择。
如果用 $z$ 路由,压缩过程可能丢失路由需要的信息。
K3 论文里也是用全宽 $x$ 做 Top-k,然后在 $\ell$ 空间计算。
\end{warningbox}
\subsection{形状总览}
\begin{center}
\begin{tabular}{llll}
\toprule
变量 & 形状 & 说明 \\
\midrule
$x$ & \shape{B, T, d} & 输入 \\
$z$ & \shape{B, T, \ell} & latent($\ell = d/2$)\\
logits & \shape{B, T, n_r} & router 输出 \\
ids & \shape{B, T, k} & Top-k 专家索引 \\
probs & \shape{B, T, k} & Top-k softmax 权重 \\
\texttt{all\_out} & \shape{n_r, B, T, \ell} & 所有 routed 专家输出 \\
$u$ & \shape{B, T, \ell} & 加权求和后的 routed 输出 \\
\texttt{shared\_out} & \shape{B, T, d} & shared 专家求和 \\
$y$ & \shape{B, T, d} & 最终输出 \\
\bottomrule
\end{tabular}
\end{center}
\subsection{本章小结}
LatentMoE 把 routed 专家限制在 $\ell = d/2$ 的 latent 空间,省参数。
SiTU-GLU 给 gate 和 up 加 $\tanh$ 软上限($\beta_1=4, \beta_2=25$),
防止低精度溢出。Shared 专家全宽,提供基础能力;routed 专家通过 Top-k 路由提供专业化能力。
+162
View File
@@ -0,0 +1,162 @@
% teach:
% gap: 读者已知各组件但不知道怎么组装成完整模型
% takeaway: K3 = Hybrid(3 KDA + 1 MLA) × DecoderBlock(attn + MoE), 末层强制 MLA
% jump: 为什么每 4 层才放一次 MLA?位置感知只需要周期性提供
% omit: 0.5b preset 的训练超参
\section{K3 混合架构}
\subsection{整体结构}
\begin{center}
\texttt{Embedding} $\to$ \texttt{DecoderBlock} $\times L$ $\to$ \texttt{RMSNorm} $\to$ \texttt{LM Head}
\end{center}
每个 \texttt{DecoderBlock} 是 Pre-Norm 残差:
\begin{lstlisting}
def forward(self, x):
x = x + self.attn(self.attn_norm(x)) # mixing
return x + self.ffn(self.ffn_norm(x)) # channel
\end{lstlisting}
\subsection{Hybrid Attention Pattern}
K3 用两种 attention 层交替:
\begin{center}
\begin{tabular}{ccccccccc}
\toprule
层 & 0 & 1 & 2 & 3 & 4 & 5 & 6 & 7 \\
\midrule
Attn & KDA & KDA & KDA & \textbf{MLA} & KDA & KDA & KDA & \textbf{MLA} \\
FFN & MoE & MoE & MoE & MoE & MoE & MoE & MoE & MoE \\
\bottomrule
\end{tabular}
\end{center}
\noindent 规则:每 4 层放 1 次 MLA(0-based 层 3, 7, 11, ...),\textbf{末层强制 MLA}。
\begin{codemathtop}{models/k3\_config.py — layer\_types}
\begin{lstlisting}
def layer_types(self) -> list[str]:
"""Hybrid: 3 KDA + 1 MLA per group, last always MLA."""
types = ["kda"] * self.num_hidden_layers
for i in range(self.num_hidden_layers):
if i % 4 == 3:
types[i] = "mla"
types[-1] = "mla" # last layer forced
return types
def layer_specs(self):
return [(kind, "moe") for kind in self.layer_types()]
\end{lstlisting}
\end{codemathtop}
\begin{knowledgebox}{为什么 KDA 不需要 RoPE?}
KDA 的 gate/decay 机制天然提供位置感知:
远的 token 衰减更多,近的保留更多。
但 softmax attention(MLA)没有这个机制,所以真实 K3 用 NoPE
(本复现的 MLA 也是 NoPE)。
位置感知从 KDA 层``渗透''到 MLA 层——3:1 的比例足够了。
\end{knowledgebox}
\subsection{CausalLM 完整数据流}
\begin{codemathtop}{models/causal\_lm.py — CausalLM}
\begin{lstlisting}
class CausalLM(nn.Module):
def __init__(self, config):
self.embedding = nn.Embedding(vocab_size, D) # [V, D]
self.blocks = ModuleList([
DecoderBlock.from_spec(config, attn, ffn)
for attn, ffn in config.layer_specs()
])
self.norm = RMSNorm(D)
self.lm_head = nn.Linear(D, vocab_size) # [V, D]
if config.tie_word_embeddings:
self.lm_head.weight = self.embedding.weight
def forward(self, input_ids, labels=None):
x = self.embedding(input_ids) # [B,T] -> [B,T,D]
if self.mixer is None: # attnres="off"
for block in self.blocks:
x = block(x) # [B,T,D] -> [B,T,D]
else:
x = self.mixer(x) # AttnRes 深度残差, 见 §9
logits = self.lm_head(self.norm(x)) # [B,T,D] -> [B,T,V]
if labels is None:
return logits
# Shifted CE: predict next token
return cross_entropy(logits[:,:-1], labels[:,1:])
\end{lstlisting}
\end{codemathtop}
\begin{knowledgebox}{残差流是可替换的}
上面的 \texttt{DecoderBlock} 逐层堆叠(\texttt{x = x + sublayer(norm(x))})
是 \texttt{config.attnres="off"} 时的默认路径。
置为 \texttt{"full"} / \texttt{"block"} 时,\texttt{CausalLM} 会把每个 block
拆成 attn / ffn 两个原子子层交给 \texttt{mixer},用\textbf{深度维注意力}
代替等权残差加法——见 \S9。真实 K3 用的是 \texttt{block} 模式。
\end{knowledgebox}
\subsection{两种配置}
\begin{center}
\begin{tabular}{lll}
\toprule
& \textbf{KDAConfig}(纯 KDA) & \textbf{K3Config}(混合) \\
\midrule
Attn & KDA only & 3 KDA + 1 MLA \\
FFN & SwiGLU & LatentMoE \\
典型规模 & \textasciitilde8M (toy) & \textasciitilde8M (toy) / \textasciitilde500M (0.5b) \\
\texttt{layer\_specs()} & \texttt{[("kda","swiglu")] * L} & \texttt{[(kind,"moe") for kind in ...]} \\
\bottomrule
\end{tabular}
\end{center}
\subsection{K3 toy 尺寸}
\begin{center}
\begin{tabular}{llll}
\toprule
参数 & 真实 K3 & toy 复现 & 缩比 \\
\midrule
$D$ & 7168 & 256 & 28$\times$ \\
$L$ & 93 & 4 & 23$\times$ \\
$H = H_V$ & 96 & 8 & 12$\times$ \\
$K = V$ & 128 & 16 & 8$\times$ \\
kv\_lora\_rank & 512 & 32 & 16$\times$ \\
q\_lora\_rank & 1536 & 64 & 24$\times$ \\
$\ell$ (MoE latent) & 3584 & 128 & 28$\times$ \\
$n_{\mathrm{routed}}$ / Top-$k$ & 896/16 & 16/2 & 56$\times$ / 8$\times$ \\
\bottomrule
\end{tabular}
\end{center}
\subsection{DecoderBlock 构建}
\begin{codemathtop}{layers/block.py — build\_attn / build\_ffn}
\begin{lstlisting}
def build_attn(config, kind: str) -> nn.Module:
if kind == "kda": return KDAAttention.from_config(config)
if kind == "mla": return GatedMLA.from_config(config)
def build_ffn(config, kind: str) -> nn.Module:
if kind == "swiglu": return SwiGLUMLP.from_config(config)
if kind == "moe": return LatentMoE.from_config(config)
class DecoderBlock(nn.Module):
def forward(self, x):
x = x + self.attn(self.attn_norm(x))
return x + self.ffn(self.ffn_norm(x))
\end{lstlisting}
\end{codemathtop}
\subsection{本章小结}
K3 架构 = Hybrid Attention(3 KDA + 1 MLA,末层强制 MLA)+ LatentMoE。
KDA 层提供线性复杂度的序列混合和位置感知(通过 decay),
MLA 层提供全局 softmax attention(NoPE,利用 KDA 渗透的位置信息)。
每层默认是 Pre-Norm 残差 DecoderBlock;\texttt{config.attnres} 可以把这条
等权残差流换成 AttnRes 深度注意力(\S9)。
+379
View File
@@ -0,0 +1,379 @@
% teach:
% gap: 读者知道 Pre-Norm 残差是"无条件等权累加", 但不知道怎么把它换成"按内容选择读哪一层"
% takeaway: AttnRes = 深度维 softmax 注意力残差; Block 版把 O(N^2) 源数压到 O(N/S); 两阶段 = inter 批量 + intra online-softmax 合并
% jump: 为什么打分用 RMSNorm 后的 v, 加权和却用原始 v
% omit: 论文里的 kernel 级调度与 pipeline 重叠
\section{Attention Residual 深度残差}
\subsection{从"等权累加"到"按内容选择"}
标准 Pre-Norm 残差把每层输出\textbf{无条件加}进残差流:
\[
x_l = x_{l-1} + f_l(x_{l-1}),
\qquad
x_N = x_0 + \sum_{l=1}^{N} f_l(x_{l-1})
\]
\noindent 展开后每一项权重恒为 1:第 3 层的输出和第 80 层的输出对最终表示的
"名义"贡献一样大,深层无法表达"我这一步应该主要读第 12 层的结果"。
AttnRes(\texttt{arXiv:2603.15031})把这个加法换成\textbf{深度维上的 softmax 注意力}:
第 $l$ 层持有一个可学习 query 向量 $w_l \in \mathbb{R}^D$,
把此前所有层的输出当成"可读的记忆":
\begin{align}
s_{l,i} &= w_l^{\top}\,\mathrm{RMS}(v_i),
& i = 0,1,\dots,l-1 \tag{A1} \\
\alpha_{l,i} &= \frac{\exp(s_{l,i})}{\sum_{j} \exp(s_{l,j})}
& \shape{n, B, T} \tag{A2} \\
h_l &= \sum_{i} \alpha_{l,i}\, v_i
& \shape{B, T, D} \tag{A3} \\
v_l &= f_l(h_l) \tag{A4}
\end{align}
\noindent 其中 $v_0 = x$(embedding 输出),$f_l$ 是已经含 Pre-Norm 的原子子层,
$\mathrm{RMS}(\cdot)$ 是不带 gain 的 RMS 归一化。
最后(\texttt{is\_final\_aggregate=True})再用一个独立 query 聚合所有源得到 $y$。
\begin{importantbox}{注意力权重是逐 token 的}
$s_{l,i}$ 的形状是 \shape{n, B, T}——每个 batch、每个位置 $t$ 都有自己的一套深度权重。
所以同一个位置在不同深度可以读不同的层,但\textbf{不跨时间混合},
因果性完全不受影响(\texttt{test\_attnres\_is\_still\_causal})。
\end{importantbox}
\subsection{DepthResidual:三个实现细节}
\begin{codemathtop}{layers/attn\_res.py — DepthResidual}
\begin{lstlisting}
class DepthResidual(nn.Module):
def __init__(self, dim, eps=1e-8, zero_init=True):
self.query = nn.Parameter(torch.zeros(dim)) # [D]
self.norm = RMSNorm(dim, eps=eps) # gain gamma
def effective_query(self):
return (self.query * self.norm.weight).float() # 折叠 gain
def forward(self, sources):
sources = stack_layers(sources) # [n,B,T,D]
q = self.effective_query() # [D]
k = rms(sources.float(), self.norm.eps) # 只用于打分
logits = einsum('d, n b t d -> n b t', q, k)
w = logits.softmax(dim=0) # 在深度维 softmax
out = einsum('n b t, n b t d -> b t d', w, sources.float())
return out.to(sources.dtype)
\end{lstlisting}
\end{codemathtop}
\paragraph{(1) gain 折叠}
RMSNorm 的可学习 gain $\gamma$ 本该作用在 key 上,但
$w^{\top}(\gamma \odot \mathrm{RMS}(v)) = (w \odot \gamma)^{\top}\mathrm{RMS}(v)$,
所以直接把 $\gamma$ 折进 query:$\tilde{w}_l = w_l \odot \gamma_l$。
少一次 \shape{n,B,T,D} 的逐元素乘法,两阶段算法里也只需要传一个向量。
\paragraph{(2) 打分用归一化的 $v$,加权和用原始 $v$}
注意 \texttt{logits} 用 \texttt{k = rms(sources)},而 \texttt{out} 用的是
\texttt{sources} 本身。
\begin{knowledgebox}{为什么这样不对称?}
打分要的是\textbf{方向}:$\mathrm{RMS}$ 之后 $s_{l,i}$ 与 $\|v_i\|$ 无关,
一层输出幅度大不会自动抢到高权重,softmax 只按"内容像不像我要读的东西"分配。\\
加权和要的是\textbf{原始信息}:如果对归一化后的 $v$ 求和,每层输出的模长
(承载着"这层贡献多大"的信息)就被抹掉了,深层的小幅修正会被放大到和主干同量级。
\end{knowledgebox}
\paragraph{(3) zero-init query}
\texttt{query} 默认初始化为 $0$ $\Rightarrow$ 所有 logits 为 $0$
$\Rightarrow$ softmax 均匀 $\Rightarrow$
\[
h_l = \frac{1}{l}\sum_{i=0}^{l-1} v_i
\]
训练起步就是\textbf{等权深度平均}(已实测:零初始化时 \texttt{forward} 输出与
\texttt{sources.mean(0)} 逐位相同),行为接近标准残差但自带 $1/l$ 缩放,
之后由梯度慢慢学出偏好。设 \texttt{zero\_init\_queries=False} 则用 $\mathcal{N}(0, 0.02^2)$。
\subsection{Full 与 Block:源数量的差别}
两种堆叠方式的区别只在\textbf{谁有资格进入源列表}:
\begin{itemize}[nosep, leftmargin=2em]
\item \texttt{FullAttnResStack}\\
保留\textbf{每一个原子层}的输出作为源,第 $l$ 层在 $l+1$ 个源上做注意力。
\item \texttt{BlockAttnResStack}\\
把 $N$ 个原子层切成大小为 $S$ 的块,\textbf{块内退化成普通求和}
(running partial $p \leftarrow p + v$),
只有\textbf{块的输出} $b_j$ 才进入源列表。
\end{itemize}
\begin{center}
\begin{tabular}{lccc}
\toprule
& \textbf{Full} & \textbf{Block ($S$)} & 标准残差 \\
\midrule
注意力源数 & 最多 $N+1$ & 最多 $N/S + 2$ & 1 \\
需保留的 \shape{B,T,D} 激活 & $O(N)$ & $O(N/S)$ & $O(1)$ \\
深度注意力 FLOPs & $O(N^2 BTD)$ & $O(N^2 BTD / S)$ & 0 \\
新增参数 & $2(N{+}1)D$ & $2(N{+}1)D$ & 0 \\
\bottomrule
\end{tabular}
\end{center}
\noindent 参数量不变(每个原子层都有自己的 query),变的是\textbf{显存与带宽}。
真实 K3($L=93$,$N=186$)取 $S=24$ 个原子层(12 个 DecoderBlock),
源数从 187 降到 $\le 10$。
\begin{codemathtop}{layers/attn\_res.py — BlockAttnResStack.forward\_naive(语义参考实现)}
\begin{lstlisting}
blocks = [x] # b_0 = embedding
partial = None
for layer_idx, (layer, residual) in enumerate(zip(self.layers, self.residuals), 1):
sources = blocks if partial is None else blocks + [partial]
h = residual(sources) # 深度注意力
out = layer(h)
partial = out if partial is None else (partial + out) # 块内: 普通累加
if (layer_idx % self.block_size == 0) or (layer_idx == len(self.layers)):
blocks.append(partial) # 块边界: 定型成一个新源
partial = None
return self.final_residual(blocks)
\end{lstlisting}
\end{codemathtop}
\subsection{两阶段算法(inter / intra)}
块内逐层跑上面的 naive 版本有个浪费:块内每一层的 query 面对的
\textbf{块间源 $b_0 \dots b_{j-1}$ 是完全相同且固定的},
唯一在变的只有 running partial $p$。于是拆成两个阶段:
\begin{enumerate}[nosep]
\item \textbf{inter(批量)}:把块内 $S$ 个 query 堆成 \shape{S, D},
对固定源做\textbf{一次}批量 einsum,拿到每个 query 的 online-softmax 三元组
$(m,\ \text{numer},\ \text{denom})$;
\item \textbf{intra(串行)}:逐层把新出现的 $p$ 作为\textbf{单个源}合并进去,
用 online softmax 的 merge 规则更新三元组,再 \texttt{normalized()} 出 $h$。
\end{enumerate}
\noindent online softmax 的三元组定义与合并规则(和 FlashAttention 同构,
只是"序列维"换成了"深度维"):
\begin{align}
m = \max_i s_i,
\quad
n = \sum_i e^{s_i - m} v_i,
\quad
d = \sum_i e^{s_i - m},
\quad
h = n / d
\tag{OS1}
\end{align}
\noindent 合并两组统计量 $(m_a, n_a, d_a)$ 与 $(m_b, n_b, d_b)$,
令 $m = \max(m_a, m_b)$、$w_a = e^{m_a - m}$、$w_b = e^{m_b - m}$:
\begin{align}
n = w_a\, n_a + w_b\, n_b,
\quad
d = w_a\, d_a + w_b\, d_b
\tag{OS2}
\end{align}
\noindent 单个源 $p$ 的三元组是 $(\,m = s_p,\ \text{numer} = p,\ \text{denom} = 1\,)$
——因为 $e^{s_p - m} = 1$,不需要真的算指数(\texttt{single\_source\_stats})。
\begin{codemathtop}{layers/attn\_res.py — \_run\_block\_two\_phase}
\begin{lstlisting}
queries = torch.stack([self.residuals[i].effective_query()
for i in range(start, end)], dim=0) # [S, D]
inter = attn_with_stats(queries, stack_layers(blocks), self.eps) # phase 1: 一次算完
partial = None
for local_idx, layer_idx in enumerate(range(start, end)): # phase 2: 串行
stats = inter.select(local_idx)
if partial is not None:
intra = single_source_stats(queries[local_idx], partial, self.eps)
stats = merge_attn_stats(stats, intra) # online softmax merge
h = stats.normalized()
out = self.layers[layer_idx](h)
partial = out if partial is None else (partial + out)
return partial
\end{lstlisting}
\end{codemathtop}
\begin{importantbox}{等价性是被测出来的,不是假设的}
\texttt{test\_block\_two\_phase\_matches\_naive} 直接对拍
\texttt{mixer.forward\_naive(emb)} 与 \texttt{mixer(emb)},
\texttt{atol=rtol=1e-5} 通过。Full 版同理:不传 \texttt{schedule\_block\_size}
走 naive,传了走两阶段,两者一致。
\end{importantbox}
\subsection{接入 CausalLM}
\subsubsection*{原子层 = 半个 DecoderBlock}
深度注意力的粒度是\textbf{原子层}而不是 DecoderBlock:
每个 block 拆成"norm + attn"和"norm + ffn"两个 Pre-Norm 原子层,
所以原子层数 $N = 2L$。
\begin{codemathtop}{models/causal\_lm.py — \_build\_mixer}
\begin{lstlisting}
atomics = []
for block in blocks:
atomics.append(BorrowedSubLayer(block.attn_norm, block.attn))
atomics.append(BorrowedSubLayer(block.ffn_norm, block.ffn))
if mode == "full":
return FullAttnResStack(D, atomics, eps=..., zero_init_queries=..., ...)
if mode == "block":
return BlockAttnResStack(D, atomics,
block_size=atomic_block_size(config.num_hidden_layers,
config.attnres_block_size), ...)
\end{lstlisting}
\end{codemathtop}
\subsubsection*{BorrowedSubLayer:借用而不注册}
\begin{codemathtop}{layers/attn\_res.py — BorrowedSubLayer}
\begin{lstlisting}
class BorrowedSubLayer(nn.Module):
def __init__(self, norm, fn):
self._borrowed = (norm, fn) # 普通 tuple, 不是 self.norm = norm
def forward(self, x):
norm, fn = self._borrowed
return fn(norm(x))
\end{lstlisting}
\end{codemathtop}
\begin{warningbox}{为什么必须用 tuple 藏起来?}
如果写成 \texttt{self.norm = norm},\texttt{nn.Module} 会把它\textbf{注册成子模块},
于是同一份权重同时挂在 \texttt{blocks.0.attn.*} 和 \texttt{mixer.layers.0.fn.*} 下:
\begin{itemize}[nosep]
\item \texttt{model.parameters()} 出现重复 $\Rightarrow$ 优化器对同一参数更新两次
\item \texttt{state\_dict()} 多出一份镜像键 $\Rightarrow$ 旧 checkpoint 加载不上
\end{itemize}
放进普通 tuple 后 \texttt{blocks.*} 仍是唯一属主,
\texttt{mixer} 下只多出 depth query 与 gain(\texttt{test\_no\_duplicate\_parameter\_ids} 守这条)。
\end{warningbox}
\subsubsection*{forward:mixer 接管整条残差流}
\begin{codemathtop}{models/causal\_lm.py — CausalLM.forward}
\begin{lstlisting}
x = self.embedding(input_ids)
if self.mixer is None:
for block in self.blocks: # attnres="off": 老路径
x = block(x)
else:
x = self.mixer(x) # full / block: DecoderBlock.forward 被完全绕过
logits = self.lm_head(self.norm(x))
\end{lstlisting}
\end{codemathtop}
\noindent 注意 \texttt{mixer} 打开后 \texttt{DecoderBlock.forward}
(\S8 里的 \texttt{x = x + attn(...)})\textbf{一次都不会被调用}——
残差加法整个交给深度注意力,DecoderBlock 退化成"两个子层的容器"。
\subsubsection*{新增参数量:可忽略}
每个 DepthResidual 只有 query \shape{D} 和 gain \shape{D},共 $N+1$ 个:
\begin{center}
\begin{tabular}{lrrr}
\toprule
配置 & $D$ / $L$ & 原子层 $N$ & 新增参数 \\
\midrule
toy (K3Config) & 256 / 4 & 8 & 4{,}608 \\
0.5b preset & 768 / 24 & 48 & 75{,}264 \\
\bottomrule
\end{tabular}
\end{center}
\subsection{配置与命令行}
\begin{center}
\begin{tabular}{lll}
\toprule
字段 & 默认 & 含义 \\
\midrule
\texttt{attnres} & \texttt{"off"} & \texttt{off} / \texttt{full} / \texttt{block} \\
\texttt{attnres\_block\_size} & \texttt{None} & 每块几个 \textbf{DecoderBlock};\texttt{None} $\to \lceil L/8 \rceil$ \\
\texttt{attnres\_zero\_init\_queries} & \texttt{True} & query 零初始化(等权起步)\\
\texttt{attnres\_final\_aggregate} & \texttt{True} & 末尾再做一次全源聚合 \\
\bottomrule
\end{tabular}
\end{center}
\begin{codemathtop}{layers/attn\_res.py — atomic\_block\_size}
\begin{lstlisting}
def atomic_block_size(num_hidden_layers, attnres_block_size):
"""DecoderBlock 数 -> 原子层数。None 时目标约 8 块。"""
layers_per_block = (attnres_block_size if attnres_block_size is not None
else max(1, (num_hidden_layers + 7) // 8))
return layers_per_block * 2 # 每个 DecoderBlock = attn|ffn 两个原子层
\end{lstlisting}
\end{codemathtop}
\noindent 单位换算是最容易踩的一处:
配置字段的单位是 \textbf{DecoderBlock 数},
而堆叠类收到的 \texttt{block\_size} 是\textbf{原子层数}($\times 2$)。
例如 $L = 24$、块大小留 \texttt{None}:
\[
\lceil 24/8 \rceil = 3 \text{ 个 DecoderBlock}
\;\to\; S = 6 \text{ 个原子层}
\;\to\; N/S = 48/6 = 8 \text{ 块}
\]
\begin{lstlisting}
uv run python train_k3.py --preset toy --attnres block --attnres-block-size 2
uv run python train_k3.py --preset 0.5b --attnres block # 块大小自动 ~L/8
\end{lstlisting}
\noindent \texttt{KDAConfig} 与 \texttt{K3Config} 都在
\texttt{\_\_post\_init\_\_} 里调 \texttt{validate\_attnres},
非法模式 / 块大小在构造时就报错。
旧 checkpoint 的 config 里没有这几个字段,加载时回落到
\texttt{off}(见 \S 9.7 验证清单最后两行)。
\subsection{验证清单}
\texttt{tests/integration/test\_attn\_res.py},14 项全过:
\begin{center}
\begin{tabular}{ll}
\toprule
测试 & 守住的性质 \\
\midrule
\texttt{default\_attnres\_is\_off} & 默认不改变任何既有行为,\texttt{mixer is None} \\
\texttt{invalid\_attnres\_rejected} & 非法 mode / \texttt{block\_size=0} 构造期报错 \\
\texttt{mixer\_kind\_and\_atomic\_count} & 原子层数 $= 2L$,块大小 $\times 2$ 换算 \\
\texttt{auto\_block\_size\_targets\_eight\_blocks} & $L=93 \to 24$(K3 $S=12$ 个 block)\\
\texttt{no\_duplicate\_parameter\_ids} & 借用不注册,参数 id / 名字均无重复 \\
\texttt{off\_and\_block\_differ\_at\_same\_seed} & 同种子下确实换了计算图 \\
\texttt{block\_two\_phase\_matches\_naive} & 两阶段 $\equiv$ naive,\texttt{atol 1e-5} \\
\texttt{attnres\_is\_still\_causal} & 改末位 token 不影响前缀 logits \\
\texttt{kda\_config\_block\_runs} & 纯 KDA 配置也能开 \\
\texttt{attnres\_ckpt\_roundtrip} & 存取后逐位一致,config 字段保真 \\
\texttt{old\_ckpt\_without\_attnres\_stays\_off} & 向后兼容 \\
\texttt{attnres\_block\_overfits\_single\_batch} & 200 步 loss $< 0.5$,能训 \\
\bottomrule
\end{tabular}
\end{center}
\begin{warningbox}{混合精度}
\texttt{DepthResidual.forward}(naive 路径)显式 \texttt{.float()} 后再算 softmax
与加权和,最后 cast 回原 dtype;两阶段路径的 query 是 fp32、源保持原 dtype,
靠 einsum 的类型提升处理。bf16 autocast 下前向实测正常。
和 \S1 的结论一致:\textbf{指数/累和一律不要放进 fp16}。
\end{warningbox}
\subsection{本章小结}
AttnRes 把残差流从"等权累加"升级成"深度维 softmax 注意力":
每层用自己的 query 决定读此前哪些层的输出,打分在 RMS 归一化后做(方向)、
加权和在原始张量上做(保留模长),query 零初始化让训练从等权平均起步。
Full 版源数随深度线性增长,Block 版把块内退化成普通求和、只让块输出进入源列表,
把源数压到 $O(N/S)$;两阶段算法进一步把块间注意力批量化,
块内用 online softmax 增量合并,与 naive 实现数值等价。
接入 \texttt{CausalLM} 时每个 DecoderBlock 拆成 attn / ffn 两个原子层,
\texttt{BorrowedSubLayer} 用普通 tuple 借用权重以免重复注册,
\texttt{attnres="off"} 保持旧路径不变。
+157
View File
@@ -0,0 +1,157 @@
% teach:
% gap: 读者会用 autograd 但不知道手写 KDA backward 的具体展开
% takeaway: backward = 逆序遍历时间步, 每步求 dq/dk/dv/dg/dbeta + 累积 dS; GVA 反传 = view+sum
% jump: 为什么 dS_dec 要加 dS_acc 和 -k⊗dr 两项
% omit: Triton backward 优化
\section{反向传播推导}
\subsection{Forward 回顾}
逐步写下 forward(省略 batch/head 下标):
\begin{align}
S_{\mathrm{dec}} &= \exp(g_t) \odot S_{t-1} \tag{F1} \\
p_t &= k_t^T S_{\mathrm{dec}} \tag{F2} \\
r_t &= v_t - p_t \tag{F3} \\
a_t &= \beta_t \cdot k_t \tag{F4} \\
S_t &= S_{\mathrm{dec}} + a_t \otimes r_t \tag{F5} \\
o_t &= q_t^T S_t \tag{F6}
\end{align}
\subsection{反传公式(BPTT,$T \to 0$)}
设 $dS_{\mathrm{acc}}$ 是从时间步 $t$ 开始累积到 $S_t$ 上的梯度。逆序遍历:
\paragraph{Step 1: $o_t = q_t^T S_t$}
\begin{align}
dS_{\mathrm{acc}} &\mathrel{+}= q_t \otimes do_t
& \xrightarrow{\texttt{einsum('bhk,bhv->bhkv')}}
& \quad \shape{B, HV, K, V} \\
dq_t &= do_t^T S_t
& \xrightarrow{\texttt{einsum('bhv,bhkv->bhk')}}
& \quad \shape{B, HV, K}
\end{align}
\paragraph{Step 2: $S_t = S_{\mathrm{dec}} + a_t \otimes r_t$}
外积的反传:$d(a \otimes r) = (\cdot)$,分解为:
\begin{align}
da_t &= \sum_v r_{t,v} \cdot dS_{\mathrm{acc},\cdot,v}
& \xrightarrow{\texttt{einsum('bhv,bhkv->bhk')}}
& \quad \shape{B, HV, K} \\
dr_t &= \sum_k a_{t,k} \cdot dS_{\mathrm{acc},k,\cdot}
& \xrightarrow{\texttt{einsum('bhk,bhkv->bhv')}}
& \quad \shape{B, HV, V}
\end{align}
\paragraph{Step 3: $a_t = \beta_t \cdot k_t$}
\begin{align}
d\beta_t &= \sum_k k_{t,k} \cdot da_{t,k}
& \xrightarrow{\texttt{einsum('bhk,bhk->bh')}}
& \quad \shape{B, HV} \\
dk_t^{(a)} &= \beta_t \cdot da_t
& & \quad \shape{B, HV, K}
\end{align}
\paragraph{Step 4: $r_t = v_t - p_t = v_t - k_t^T S_{\mathrm{dec}}$}
\begin{align}
dv_t &= dr_t & & \shape{B, HV, V} \\
dk_t^{(r)} &= -S_{\mathrm{dec}}^T \cdot dr_t
& \xrightarrow{\texttt{einsum('bhv,bhkv->bhk')}}
& \quad \shape{B, HV, K} \\
dS_{\mathrm{dec}}^{(r)} &= -k_t \otimes dr_t
& \xrightarrow{\texttt{einsum('bhv,bhk->bhkv')}}
& \quad \shape{B, HV, K, V}
\end{align}
\paragraph{Step 5: 合并 $dS_{\mathrm{dec}}$ 并传递 $dg_t$, $dS_{t-1}$}
\[
dS_{\mathrm{dec}}^{\mathrm{total}} = dS_{\mathrm{acc}} + dS_{\mathrm{dec}}^{(r)}
= dS_{\mathrm{acc}} - k_t \otimes dr_t
\]
因为 $S_{\mathrm{dec}} = \exp(g_t) \odot S_{t-1}$:
\begin{align}
dg_t &= S_{\mathrm{dec}} \odot dS_{\mathrm{dec}}^{\mathrm{total}}
& \xrightarrow{\texttt{einsum('bhkv,bhkv->bhk')}}
& \quad \shape{B, HV, K} \\
dS_{t-1} &= \exp(g_t) \odot dS_{\mathrm{dec}}^{\mathrm{total}}
& & \quad \shape{B, HV, K, V}
\end{align}
\paragraph{Step 6: 合并 $dk_t$ 和 GVA 归约}
\[
dk_t = dk_t^{(a)} + dk_t^{(r)}
= \beta_t \cdot da_t - S_{\mathrm{dec}}^T \cdot dr_t
\]
GVA 反传($H_V \to H$):
\[
dq_H = dq_{H_V}.\texttt{view}(B, T, H, G, K).\texttt{sum}(\text{dim}=3) \cdot \mathrm{scale}
\]
\[
dk_H = dk_{H_V}.\texttt{view}(B, T, H, G, K).\texttt{sum}(\text{dim}=3)
\]
\subsection{代码对照}
\begin{codemathtop}{ops/reference/recurrent.py — KDAFunction.backward}
\begin{lstlisting}
for t in range(T - 1, -1, -1):
q_t, k_t, b_t = q_ts[:,t], k_ts[:,t], b_ts[:,t]
S_dec, r_t, a_t = S_decs[:,t], r_ts[:,t], a_ts[:,t]
exp_g_t, do_t = exp_g_ts[:,t], do[:,t]
# Step 1: o_t = q_t . S_t
S_t = S_dec + einsum('bhk,bhv->bhkv', a_t, r_t)
dS_acc += einsum('bhk,bhv->bhkv', q_t, do_t)
dq_e[:,t] = einsum('bhv,bhkv->bhk', do_t, S_t)
# Step 2: outer product grads
da_t = einsum('bhv,bhkv->bhk', r_t, dS_acc)
dr_t = einsum('bhk,bhkv->bhv', a_t, dS_acc)
# Step 3: a_t = beta_t * k_t
dbeta[:,t] = einsum('bhk,bhk->bh', k_t, da_t)
dk_t_a = b_t.unsqueeze(-1) * da_t
# Step 4: r_t = v_t - k_t . S_dec
dv[:,t] = dr_t
dS_dec_from_r = -einsum('bhv,bhk->bhkv', dr_t, k_t)
dk_t_r = -einsum('bhv,bhkv->bhk', dr_t, S_dec)
# Step 5: S_dec = exp(g) * S_{t-1}
dS_dec_total = dS_acc + dS_dec_from_r
dk_e[:,t] = dk_t_a + dk_t_r
dg[:,t] = einsum('bhkv,bhkv->bhk', S_dec, dS_dec_total)
dS_acc = exp_g_t.unsqueeze(-1) * dS_dec_total
# Step 6: GVA reduce
dq_H = dq_e.view(B,T,H,G,K).sum(dim=3) * scale
dk_H = dk_e.view(B,T,H,G,K).sum(dim=3)
\end{lstlisting}
\end{codemathtop}
\begin{warningbox}{$dS_{\mathrm{dec}}^{\mathrm{total}}$ 为什么包含两项?}
$S_t = S_{\mathrm{dec}} + a_t \otimes r_t$,$S_{\mathrm{dec}}$ 同时参与了:
\begin{enumerate}[nosep]
\item 直接传递到 $dS_{\mathrm{acc}}$(作为 $S_t$ 的一部分被读出)
\item 通过 $r_t = v_t - k_t \cdot S_{\mathrm{dec}}$ 参与 delta rule
\end{enumerate}
所以 $dS_{\mathrm{dec}}^{\mathrm{total}} = dS_{\mathrm{acc}} + dS_{\mathrm{dec}}^{(r)}$,
两条路径的梯度要\textbf{加}起来(chain rule 分叉处求和)。
\end{warningbox}
\subsection{本章小结}
KDA backward 是 BPTT 展开:逆序遍历时间步,每步 6 个 einsum + 一次 $dS$ 累积更新。
GVA 反传在最后做 \texttt{view+sum}。手写 backward 的关键是正确处理
$dS_{\mathrm{dec}}$ 的两条梯度路径(直接传递 + 通过 $r_t$ 的 delta rule 路径)。
+183
View File
@@ -0,0 +1,183 @@
% teach:
% gap: none — this is a reference appendix
% takeaway: 一表查所有符号
% jump: none
% omit: none
\section{符号表}
\subsection{形状参数}
\begin{center}
\begin{tabular}{lll}
\toprule
符号 & 含义 & 典型值 (toy) \\
\midrule
$B$ & batch size & 2--4 \\
$T$ & 序列长度 & 128--2048 \\
$D$ & hidden\_size & 64 / 256 \\
$H$ & query/key 头数 & 4 / 8 \\
$H_V$ & value 头数(GVA) & $G \cdot H$ \\
$G$ & GVA 组数 & $H_V / H$ \\
$K$ & key/query 头维度 & 16 \\
$V$ & value 头维度($= K$) & 16 \\
$C$ & chunk\_size & 16 / 64 \\
$r$ & KV latent rank (MLA) & 32 \\
$d_q$ & MLA query head dim & 16 \\
$d_v$ & MLA value head dim & 16 \\
$\ell$ & MoE latent width ($= D/2$) & 128 \\
$n_r$ & routed 专家数 & 16 \\
$k$ & Top-$k$ & 2 \\
$n_s$ & shared 专家数 & 2 \\
$d_{\mathrm{ff}}$ & 专家中间维度 & 96 \\
$N$ & AttnRes 原子层数 ($= 2L$) & 8 \\
$S$ & AttnRes 块大小(原子层) & 2--24 \\
\bottomrule
\end{tabular}
\end{center}
\subsection{KDA 变量}
\begin{center}
\begin{tabular}{llp{7cm}}
\toprule
符号 & 形状 & 含义 \\
\midrule
$q_t$ & \shape{B, HV, K} & query(已 GVA 展开 + scale) \\
$k_t$ & \shape{B, HV, K} & key(已 GVA 展开) \\
$v_t$ & \shape{B, HV, V} & value \\
$g_t$ & \shape{B, HV, K} & gate(log-space 衰减,逐维逐头) \\
$\beta_t$ & \shape{B, HV} & 写入强度 \\
$S_t$ & \shape{B, HV, K, V} & KV 状态矩阵 \\
$S_{\mathrm{dec}}$ & \shape{B, HV, K, V} & 衰减后的状态 \\
$p_t$ & \shape{B, HV, V} & 旧状态对 $k_t$ 的预测 \\
$r_t$ & \shape{B, HV, V} & delta rule 残差 = $v_t - p_t$ \\
$a_t$ & \shape{B, HV, K} & 写入向量 = $\beta_t \cdot k_t$ \\
$o_t$ & \shape{B, HV, V} & 读出 = $q_t \cdot S_t$ \\
\bottomrule
\end{tabular}
\end{center}
\subsection{Gate 变量}
\begin{center}
\begin{tabular}{llp{6cm}}
\toprule
符号 & 形状 & 含义 \\
\midrule
$g_{\mathrm{raw}}$ & \shape{B, T, HV, K} & gate 投影原始输出 \\
$A_{\log}$ & \shape{HV} & head-wise 衰减参数 (log-space) \\
$\Delta_b$ & \shape{HV, K} & per-dim gate bias \\
$\mathrm{rate}$ & \shape{HV, 1} & $\exp(A_{\log})$ \\
$\mathrm{input}$ & \shape{B, T, HV, K} & $g_{\mathrm{raw}} + \Delta_b$ \\
$L$ & 标量 & lower\_bound ($-5.0$) \\
\bottomrule
\end{tabular}
\end{center}
\subsection{MLA 变量}
\begin{center}
\begin{tabular}{llp{6cm}}
\toprule
符号 & 形状 & 含义 \\
\midrule
$c$ & \shape{B, T, r} & KV latent(推理时缓存这个) \\
$q$ & \shape{B, T, H, d_q} & query(低秩路径输出) \\
$W_{UK}$ & \shape{H, d_q, r} & key 解压矩阵(吸收进 $q$) \\
$W_{UV}$ & \shape{H, d_v, r} & value 解压矩阵 \\
$q_{\mathrm{abs}}$ & \shape{B, T, H, r} & 吸收后的 query \\
score & \shape{B, H, T, T} & $q_{\mathrm{abs}} \cdot c^T$ \\
attn & \shape{B, H, T, T} & causal softmax \\
$\tilde{o}_{\mathrm{lat}}$ & \shape{B, H, T, r} & latent 加权输出 \\
$\tilde{o}$ & \shape{B, H, T, d_v} & 解压后的输出 \\
gate & \shape{B, T, H \cdot d_v} & $\sigma(W_g x)$ \\
\bottomrule
\end{tabular}
\end{center}
\subsection{LatentMoE 变量}
\begin{center}
\begin{tabular}{llp{6cm}}
\toprule
符号 & 形状 & 含义 \\
\midrule
$x$ & \shape{B, T, D} & 输入 \\
$z$ & \shape{B, T, \ell} & latent ($\ell = D/2$) \\
logits & \shape{B, T, n_r} & router logits \\
ids & \shape{B, T, k} & Top-$k$ 专家索引 \\
probs & \shape{B, T, k} & softmax 权重 \\
$u$ & \shape{B, T, \ell} & routed 加权输出 \\
$s$ & \shape{B, T, D} & shared 专家求和 \\
$y$ & \shape{B, T, D} & $s + W_\uparrow \mathrm{RMSNorm}(u)$ \\
\bottomrule
\end{tabular}
\end{center}
\subsection{AttnRes 变量}
\begin{center}
\begin{tabular}{llp{6.4cm}}
\toprule
符号 & 形状 & 含义 \\
\midrule
$v_i$ & \shape{B, T, D} & 第 $i$ 个源($v_0 =$ embedding 输出) \\
$w_l$ & \shape{D} & 第 $l$ 层的 depth query(零初始化) \\
$\gamma_l$ & \shape{D} & DepthResidual 的 RMSNorm gain \\
$\tilde{w}_l$ & \shape{D} & 折叠后的 query $= w_l \odot \gamma_l$ \\
$s_{l,i}$ & \shape{n, B, T} & 深度打分 $= \tilde{w}_l^{\top}\mathrm{RMS}(v_i)$ \\
$\alpha_{l,i}$ & \shape{n, B, T} & 深度维 softmax 权重 \\
$h_l$ & \shape{B, T, D} & 第 $l$ 层的输入 $= \sum_i \alpha_{l,i} v_i$ \\
$b_j$ & \shape{B, T, D} & 第 $j$ 个块的输出(Block 版的源) \\
$p$ & \shape{B, T, D} & 块内 running partial \\
$m, n, d$ & \shape{B, T} / \shape{B,T,D} / \shape{B,T} & online softmax 三元组 \\
\bottomrule
\end{tabular}
\end{center}
\subsection{Einsum 速查}
\begin{center}
\small
\begin{tabular}{p{6cm}lp{3.5cm}}
\toprule
操作 & einsum & 结果形状 \\
\midrule
key 查状态 & \texttt{'bhk,bhkv->bhv'} & $p_t$ \shape{B,HV,V} \\
外积写入 & \texttt{'bhk,bhv->bhkv'} & $a_t \otimes r_t$ \shape{B,HV,K,V} \\
读出 & \texttt{'bhk,bhkv->bhv'} & $o_t$ \shape{B,HV,V} \\
MLA 吸收 $W_{UK}$ & \texttt{'bthd,hdj->bthj'} & $q_{\mathrm{abs}}$ \shape{B,T,H,r} \\
MLA 打分 & \texttt{'bthj,bsj->bhts'} & score \shape{B,H,T,T} \\
MLA latent 加权 & \texttt{'bhts,bsj->bhtj'} & $\tilde{o}_{\mathrm{lat}}$ \shape{B,H,T,r} \\
MLA 解压 & \texttt{'bhtj,hvj->bhtv'} & $\tilde{o}$ \shape{B,H,T,d_v} \\
AttnRes 深度打分 & \texttt{'d,nbtd->nbt'} & $s_{l,i}$ \shape{n,B,T} \\
AttnRes 深度加权和 & \texttt{'nbt,nbtd->btd'} & $h_l$ \shape{B,T,D} \\
AttnRes 批量打分(inter) & \texttt{'qd,nbtd->qnbt'} & logits \shape{S,n,B,T} \\
\bottomrule
\end{tabular}
\end{center}
\subsection{总结与延伸}
\subsubsection*{核心要点回顾}
\begin{enumerate}[nosep]
\item \textbf{KDA} = delta rule 状态更新 + gate 衰减,线性复杂度
\item \textbf{分块} = chunk 内下三角解 + chunk 间状态递推,等价于 naive recurrent
\item \textbf{GVA} = $H_V = G \cdot H$,forward repeat\_interleave / backward view+sum
\item \textbf{MLA} = 低秩 latent + 矩阵吸收,KV cache 从 $2Hd$ 降到 $r$
\item \textbf{LatentMoE} = shared 全宽 + routed 半宽 latent + SiTU-GLU 防溢出
\item \textbf{K3 Hybrid} = 3 KDA + 1 MLA,KDA 提供位置感知
\item \textbf{AttnRes} = 深度维 softmax 残差,Block 版把源数压到 $O(N/S)$,
两阶段 = inter 批量 + intra online-softmax 合并
\end{enumerate}
\subsubsection*{未完成项}
\begin{itemize}[nosep]
\item L5 — 项目内自研 fused gate Triton kernel
\item L6 — recurrent decode cache(推理加速)
\item AttnRes 与 recurrent decode 的组合(增量解码时的深度源缓存)
\item AttnRes 开 / 关的收敛质量对比实验(目前只验证了等价性与可训练性)
\end{itemize}