Back

Attention Residuals 论文深度解读:Kimi K3 的核心架构创新

论文信息

这篇论文提出的 Attention Residuals(AttnRes)是 Kimi K3 的核心架构创新之一。理解这篇论文,就能理解 Kimi K3 为什么能在 2.8 万亿参数规模下依然保持高效训练和推理。


一、问题背景:残差连接的"原罪"

1.1 标准残差连接

现代大语言模型(LLM)都使用残差连接(Residual Connection)。简单来说,每一层的输出会"跳过"自己,直接加到下一层的输入上:

h_l = h_{l-1} + f(h_{l-1})

其中 h_l 是第 l 层的隐藏状态,f 是该层的变换函数。

这个设计的初衷是让梯度能够直接回传到任意层,解决了深度网络训练困难的问题。

1.2 PreNorm 的稀释问题

当前主流架构使用 PreNorm(先归一化再变换):

h_l = h_{l-1} + f(Norm(h_{l-1}))

但 PreNorm 有一个严重问题:隐藏状态的范数会随着深度线性增长

展开残差连接,第 l 层的输入实际上是所有前序层输出的简单求和:

h_l = h_1 + f_1(h_1) + f_2(h_2) + ... + f_{l-1}(h_{l-1})

这意味着:

  • 隐藏状态范数按 O(L) 增长
  • 每一层的相对贡献被不断稀释
  • 深层信息被淹没,无法被选择性检索

论文指出:经验上,相当大比例的层可以被剪枝而几乎不损失性能——这说明很多层在"白干"。

1.3 现有改进方案的局限

方案 思路 局限
Highway Networks 引入门控机制 每层仍只能访问前一层
DenseFormer 允许访问所有前序层 使用固定权重,无输入依赖性
mHC 多流并行 复杂度高,难以扩展

核心问题:这些方案都没有解决"选择性访问"的问题——模型无法根据输入内容动态决定从哪些层获取信息。


二、核心洞察:时间与深度的对偶性

论文提出了一个关键洞察:深度维度的信息累积,与序列维度的 RNN 递推,在数学形式上是完全对偶的

2.1 序列维度的演进

阶段 方法 核心思想
RNN 循环递推 压缩所有历史到一个状态
Attention 注意力机制 允许选择性访问所有历史位置

2.2 深度维度的演进(本文贡献)

阶段 方法 核心思想
残差连接 固定累加 压缩所有前序层到一个状态
AttnRes 深度注意力 允许选择性访问所有前序层

论文的核心贡献:将 Attention 的思想从序列维度迁移到深度维度,提出 Attention Residuals。


三、Attention Residuals 详解

3.1 Full AttnRes:完整版本

核心公式

h_l = Σ α_{i→l} · v_i

其中:

  • v_0 = h_1(词嵌入)
  • v_i = f_i(h_i)(第 i 层的输出)
  • α_{i→l} 是 softmax 注意力权重

注意力权重计算

α_{i→l} = φ(w_l, k_i) / Σ φ(w_l, k_j)
φ(q, k) = exp(q^T · RMSNorm(k))

关键设计

  • 每层有一个可学习的伪查询向量 w_l ∈ R^d
  • 键和值都是前序层的输出
  • 使用 RMSNorm 防止大范数层主导注意力权重

3.2 复杂度分析

维度 Full AttnRes 标准残差
计算量 O(L²d) O(Ld)
内存 O(Ld) O(d)

由于网络深度 L 远小于序列长度(通常 L < 1000),O(L²) 的计算量是可接受的。

3.3 Block AttnRes:实用版本

Full AttnRes 在大规模训练中有问题:

  • 流水线并行需要跨阶段传输所有层输出
  • 内存和通信开销按 O(Ld) 增长

解决方案:Block AttnRes

将 L 层分成 N 个块(block)
- 块内:使用标准残差累加
- 块间:使用 AttnRes 进行选择性聚合

效果

  • 内存从 O(Ld) 降到 O(Nd)
  • 计算从 O(L²) 降到 O(N²)
  • 经验上 N ≈ 8 就能恢复大部分收益

四、工程优化:让 AttnRes 真正可用

4.1 训练优化:跨阶段缓存

在流水线并行中,Block AttnRes 需要跨阶段传输块表示。

朴素方法:每次转换都传输所有累积的块 → 冗余通信

优化方法:跨阶段缓存

  • 每个物理阶段缓存之前接收的块
  • 后续虚拟阶段只传输增量块

效果:峰值通信成本从 O(C) 降到 O(P),提升 V 倍。

4.2 推理优化:两阶段计算

Block AttnRes 的推理需要:

  1. 跨块注意力(并行)
  2. 块内注意力(顺序)

两阶段策略

  • 阶段1:并行计算跨块注意力
  • 阶段2:顺序处理块内注意力,使用 online softmax 合并结果

效果:推理延迟增加 < 2%。


五、实验结果

5.1 Scaling Law 实验

在固定计算预算下,AttnRes 始终优于基线:

模型 基线 Loss AttnRes Loss 提升
小模型 1.766 1.737 0.029
中模型 1.747 1.720 0.027
大模型 1.730 1.705 0.025

关键发现:Block AttnRes (N=8) 匹配了使用 1.25 倍计算量训练的基线模型。

5.2 下游任务性能

在 Kimi Linear 架构(48B 总参数 / 3B 激活参数)上预训练 1.4T tokens:

任务类别 基线 AttnRes 提升
通用理解
MMLU 73.5 74.6 +1.1
GPQA-Diamond 36.9 44.4 +7.5
BBH 76.3 78.0 +1.7
数学与代码
GSM8K 81.7 82.4 +0.7
Math 53.5 57.1 +3.6
HumanEval 59.1 62.2 +3.1
MBPP 72.0 73.9 +1.9
中文理解
CMMLU 82.0 82.9 +0.9
C-Eval 79.6 82.5 +2.9

关键发现

  • 多步推理任务提升最明显(GPQA-Diamond +7.5)
  • 代码生成提升显著(HumanEval +3.1)
  • 所有任务都有提升,没有退化

5.3 训练动态分析

AttnRes 解决了 PreNorm 稀释问题:

指标 基线 AttnRes
隐藏状态范数 随深度单调增长 有界,呈周期性模式
梯度分布 早期层梯度 disproportionately 大 更均匀分布

六、消融实验:关键设计选择

6.1 跨层访问的重要性

方法 Loss 说明
基线 (PreNorm) 1.766 只能访问前一层
滑动窗口 (W=8) 1.764 只访问最近 8 层
Block AttnRes 1.746 块级访问
Full AttnRes 1.737 访问所有前序层

结论:选择性访问远距离层比访问多个近邻层更重要。

6.2 注意力机制组件

变体 Loss 说明
Full AttnRes 1.737 完整版本
输入无关混合 1.749 移除 query/key
使用 sigmoid 1.741 替换 softmax
移除 RMSNorm 1.743 不归一化 key
多头注意力 1.752 H=16

关键发现

  • 输入依赖的查询至关重要(移除后 loss 上升 0.012)
  • softmax 优于 sigmoid(竞争归一化带来更尖锐的选择)
  • RMSNorm 必不可少(防止大范数层主导)
  • 单头优于多头(最优深度混合在通道间基本一致)

6.3 最优架构偏好

在固定计算和参数预算下,AttnRes 改变了最优架构:

方法 最优 d_model/L_b 含义
基线 ~60 较浅较宽
AttnRes ~45 更深更窄

结论:AttnRes 能更有效地利用深度,在相同参数预算下倾向于更深的网络。


七、可视化:学到的注意力模式

论文可视化了 16 层模型的深度注意力权重分布,发现三个关键模式:

7.1 保持局部性

每层仍然主要关注其直接前驱(对角线主导),但出现了选择性的非对角集中——学习到的"跳跃连接"。

7.2 层特化

  • 词嵌入在整个网络中保持非平凡权重(尤其是注意力层前)
  • 注意力层前:保持较广的感受野
  • MLP 层前:更依赖最近的表示

7.3 Block AttnRes 保留结构

块级压缩充当隐式正则化,保留了完整版本的关键信息通路。


八、与 Kimi K3 的关系

Attention Residuals 是 Kimi K3 的两大核心架构创新之一(另一个是 KDA 混合线性注意力)。

8.1 在 K3 中的应用

组件 作用
KDA 序列维度的线性注意力,降低推理成本
AttnRes 深度维度的选择性聚合,解决 PreNorm 稀释

两者结合,使 K3 能够在 2.8 万亿参数规模下:

  • 保持训练稳定性
  • 实现高效推理
  • 获得更好的性能

8.2 在 K3 中的配置

根据 K3 官方博客:

  • 使用 Block AttnRes(实用版本)
  • 块数 N ≈ 8(经验最优)
  • 推理延迟增加 < 2%

九、总结与启示

9.1 核心贡献

  1. 理论洞察:揭示了时间-深度对偶性,将 Attention 思想迁移到深度维度
  2. 方法创新:提出 AttnRes,用 softmax 注意力替代固定残差累加
  3. 工程优化:Block AttnRes + 跨阶段缓存 + 两阶段推理,使方法真正可用
  4. 实验验证:在 scaling law、下游任务、训练动态上全面验证

9.2 对行业的启示

启示 说明
架构创新仍有空间 不是只有数据和算力,架构设计同样重要
对偶思维 序列维度的成功方法可以迁移到深度维度
工程与理论并重 好的方法必须考虑大规模训练的可行性
开源价值 论文和代码都开源,推动社区进步

9.3 未来方向

论文提到的未来工作:

  • 更细粒度的块大小或 Full AttnRes(随着硬件改进)
  • 结合更高效的线性复杂度替代方案
  • 探索其他维度上的类似对偶性

附录:关键公式汇总

A.1 标准残差连接

h_l = h_{l-1} + f(h_{l-1})

A.2 Full AttnRes

h_l = Σ_{i=0}^{l-1} α_{i→l} · v_i
α_{i→l} = exp(w_l^T · RMSNorm(k_i)) / Σ_j exp(w_l^T · RMSNorm(k_j))

A.3 Block AttnRes

块内:b_n = Σ_{j∈B_n} f_j(h_j)
块间:h_l = Σ_{n=0}^{N-1} α_{n→l} · b_n + α_{n'→l} · b_{n'}^{i-1}

本文基于 arXiv:2603.15031 论文翻译解读,如有错误欢迎指正。

评论