Kimi K3 模型结构(4):深度维度,Attention Residuals

Kimi K3 模型结构(4):深度维度,Attention Residuals

💡 原文中文,约16800字,阅读约需40分钟。
📝

内容提要

本文解析Kimi K3模型中的Attention Residuals(注意力残差)结构。该机制将残差连接视为深度方向上的RNN,通过可学习伪query对embedding及各层输出做softmax加权,替代固定求和。K3采用Block形式,将93层分为8块,块内普通求和,块间注意力聚合,共9个来源。相比普通残差,此设计提升多步推理能力,降低输出幅度增长,并支持高效两阶段推理。

🔎

延伸解读

残差连接为何要升级

普通残差连接在深度方向上存在三个问题:权重固定、信息丢失后无法找回、输出幅度随深度增长。作者将其类比为深度方向上的RNN,只能通过前一个状态间接接触历史。Attention Residuals正是借鉴序列方向上的注意力机制,让每一层都能带权重地访问全部历史输出,从而解决这些问题。

Block形式的设计权衡

Full形式虽然理论上更灵活,但需要保留所有层的输出,内存和通信开销大。Block形式通过将层分组,块内普通求和,块间注意力聚合,把内存从O(Ld)降到O(Nd)。K3选择8块,是因为论文扫描显示N≈8已能获得绝大部分收益,且块大小扫描表明继续增加块数收益有限。

推理效率的关键设计

伪query w_l是参数而非投影,使得块间打分可以提前批量计算,与块内计算重叠。两阶段推理配合online softmax,将额外访存控制在每token每层约5.5d,相比Full形式的24d大幅降低,实测推理延迟开销小于2%。这是Block形式能落地的关键。

术语与实现细节提醒

论文中'layer'指子层,一个decoder层贡献两个子层,因此K3的93层对应186个子层。代码中attn_res_block_size=12是按decoder层计,每个decoder层内实际有两次AttnRes(注意力前和MoE前)。理解这些口径差异,有助于正确对照论文与代码。

Q&A

Kimi K3 中的 Attention Residuals 是什么?

Attention Residuals 是一种将残差连接视为深度方向上的 RNN 的机制,它用可学习的伪 query 对 embedding 和各层输出做 softmax 加权,替代了固定的求和。K3 采用 Block 形式,将 93 层分为 8 块,块内普通求和,块间注意力聚合,共 9 个来源。

普通残差连接有哪些问题?

普通残差连接有三个问题:1. 权重固定,所有层拿到相同的和,无法按需组合;2. 信息被求和糊掉,后面的层无法单独取回某一层的输出;3. 输出幅度随深度增长,导致深层训练不稳定。

Full AttnRes 和 Block AttnRes 有什么区别?

Full AttnRes 中每个子层的输出都单独保留,任何后续子层都能单独取回它;Block AttnRes 将子层分成块,块内输出先求和,只有块的和能被跨块取回。Block 形式降低了内存和通信开销,从 O(Ld) 降到 O(Nd)。

Kimi K3 中 Attention Residuals 是如何分块的?

K3 将 93 层 decoder layer 按每 12 层分为一块,共 7 个满块和 1 个 9 层的尾块,加上 embedding 作为第 0 块,共 9 个来源。块内普通求和,块间用注意力聚合。

Attention Residuals 中的伪 query 是什么?为什么不用投影?

伪 query 是每个子层各自拥有的一个可学习向量,所有 token 和位置共用,不依赖输入。论文消融过用隐藏状态投影的 query,loss 更好,但推理时被迫顺序访存,所以放弃。使用参数化的 query 可以在块内所有层运行前批量算好打分,支持高效的两阶段推理。

Attention Residuals 带来了哪些收益?

收益包括:1. 提升多步推理能力,如 GPQA-Diamond 提升 7.5 分;2. 降低输出幅度增长,使训练更稳定;3. 梯度分布更均匀;4. 学到的权重模式显示跨层 skip connection 和 attention sink;5. 模型形状偏好更深的模型。

Attention Residuals 与 Hyper-Connections 有何区别?

Attention Residuals 是深度上的 softmax 注意力,秩为 L;Hyper-Connections (mHC) 是深度上的线性注意力,状态是矩阵。两者关系类似序列方向上线性注意力和 softmax 注意力的关系。

Kimi K3 推理时如何高效计算 Attention Residuals?

推理时采用两阶段:阶段一并行计算所有子层对已完成块的注意力,因为伪 query 是参数,可提前算好;阶段二顺序处理块内新增的部分和,用 online softmax 合并。这样减少了访存,延迟开销小于 2%。

🏷️

标签

➡️

继续阅读