Files
K3/notes/sections/sec-09.tex
T
dela 584f7e9e73 Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias
backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
2026-08-25 14:43:17 +08:00

380 lines
16 KiB
TeX
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
% teach:
% gap: 读者知道 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"} 保持旧路径不变。