% 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"} 保持旧路径不变。