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,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。
|
||||
@@ -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\%。
|
||||
@@ -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 逐位相同。
|
||||
@@ -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$ 维投影。
|
||||
@@ -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 内部完成。
|
||||
@@ -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。
|
||||
@@ -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 路由提供专业化能力。
|
||||
@@ -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)。
|
||||
@@ -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"} 保持旧路径不变。
|
||||
@@ -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 路径)。
|
||||
@@ -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}
|
||||
Reference in New Issue
Block a user