Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
380 lines
16 KiB
TeX
380 lines
16 KiB
TeX
% 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"} 保持旧路径不变。
|