arXiv'26 | Kimi K3 的 AttnRes 与 Stable LatentMoE:深度和宽度怎么一起扩展

arXiv’26 | Kimi K3 的 AttnRes 与 Stable LatentMoE:深度和宽度怎么一起扩展

原文:Kimi K3: Open Frontier Intelligence

这是 Kimi K3 系列的第二篇。上一篇讲了序列维度上的 KDA:历史 token 怎么压缩、怎么遗忘、怎么在 GPU 上并行。本篇转到另外两个维度:层与层之间的信息怎么传,专家与专家之间的容量怎么扩展。

系列导航:导读上一篇:KDA 与 Hybrid Attention下一篇:训练、Agent 与 Serving Infra

K3 的配置很容易让人先记住一个大数字:2.78T 总参数、104.2B 激活参数、896 个 routed expert 每个 token 选 16 个。但真正值得理解的不是数字本身,而是两个问题:层数变深以后,为什么上一层残差不一定够用;专家变多以后,为什么 router 和激活值会先不稳定。

1. 深度方向的默认残差有什么限制

标准 Transformer 每层通常做这样的事情:

\[h_l=h_{l-1}+F_l(h_{l-1})\]

$h_{l-1}$ 是上一层的 hidden state,$F_l$ 是第 $l$ 层的 attention 或 FFN 变换。这个结构的优点是简单,梯度也容易沿着残差往回传;但它隐含了一个很强的限制:第 $l$ 层只能直接拿到第 $l-1$ 层的结果。

假设一个 90 层模型的第 80 层需要某个早期层提取的局部语法特征。这个特征只能混在第 10 层到第 79 层的连续变换里传过来,中间每一层都可能改变它。模型当然可以学着把信息藏在残差里,但这条通路不是显式的“回到第 10 层查一下”。

这和时间方向的 recurrent state 有点像:历史信息被压进一个状态,后面只能使用这个被反复更新的状态。K3 的 AttnRes(Attention Residuals)想做的是,把“只接上一层”改成“从前面多个层级中检索”。

2. AttnRes:把残差连接改成深度上的 attention

假设当前是第 $l$ 层,前面已经产生了 $l$ 个表示:

\[H_{<l}=[h_0,h_1,\ldots,h_{l-1}]\in\mathbb R^{l\times d}\]

其中 $h_0$ 可以是 token embedding,$d$ 是 hidden dimension。AttnRes 为当前层准备一个可学习的 pseudo query $\bm w_l\in\mathbb R^d$,把前面每层的输出当作 key/value:

# H_prev: [l, d],前面 l 个层级的表示
# w_l: [d],第 l 层自己的可学习 query
keys = rmsnorm(H_prev)                 # [l, d]
weights = softmax(w_l @ keys.T, dim=-1) # [l]
residual = weights @ H_prev             # [d]

公式写出来是:

\[\bm a_l=\operatorname{softmax}\left(\bm w_l\,\operatorname{RMSNorm}(H_{<l})^\top\right), \qquad r_l=\bm a_lH_{<l}\]

然后第 $l$ 层不再只接收 $h_{l-1}$,而是接收检索结果 $r_l$ 与当前层变换的组合。这里的 attention 没有发生在 token 之间,而是发生在层的深度轴上。

为什么必须先做 RMSNorm?不同层的输出幅值可能差别很大。如果直接点积,某一层仅仅因为数值范数大,就可能在 softmax 中占据绝大多数权重;RMSNorm 把比较重点放回方向和内容,而不是谁的向量更长。

2.1 用 4 层的例子走一遍

假设前面有 embedding、layer 1、layer 2、layer 3 四个表示,当前 layer 4 的 pseudo query 计算出:

attention over depth = [0.05, 0.10, 0.70, 0.15]

那么 layer 4 使用的 residual 不是 h_3,而是:

\[r_4=0.05h_0+0.10h_1+0.70h_2+0.15h_3\]

如果当前任务需要 layer 2 的特征,模型可以直接把 0.70 的权重放在那里;如果另一个 token 需要最新的局部信息,权重也可以移动到 $h_3$。这比所有层都无差别相加多了一层选择能力。

3. 为什么 K3 使用 Block AttnRes

Full AttnRes 的问题是显存和通信。每一层都要保留前面所有层的输出,若有 $L$ 层、hidden dimension 为 $d$,需要保存的表示规模大致是 $O(Ld)$。模型跨 GPU 做 pipeline parallelism 时,这些表示也会产生额外通信。

K3 的解决办法是把层分成 block。以 93 层为例,可以切成大约 8 个 block,每个 block 内的层输出先汇总成一个 block representation;再加上 embedding,一共约 9 个可检索的来源。每层只对这 9 个表示做 attention,而不是对全部 93 层做 attention。

embedding  ───────────────┐
layers 1-12  -> block 1 ──┤
layers 13-24 -> block 2 ──┤──> 当前层的 depth attention
...                       │
layers 85-93 -> block 8 ──┘

这个近似的直觉是:同一个 block 内相邻层的表示往往比较相近,没必要每一层都单独作为 key/value;但跨 block 的信息差异更大,值得保留独立入口。论文的消融显示,block 数约为 8 时可以获得 Full 版本的大部分收益。

所以 Block AttnRes 不是简单地“把 full 版本砍小”,而是在两个开销之间取平衡:保留多个深度层级的可检索性,同时把持久状态从 93 份压缩到约 9 份。

Kimi K3 的三维架构

上图把 K3 的三个方向放在一起:token mixing 对应 KDA/MLA,layer mixing 对应 AttnRes,channel mixing 对应 MoE。接下来进入宽度方向。

4. 896 个专家为什么要先进入 latent space

MoE(Mixture-of-Experts,混合专家)把 FFN 复制成很多专家,再由 router 为每个 token 选择少数几个。总参数可以很大,但每个 token 只激活一小部分。

K3 使用 896 个 routed expert,每个 token 选择 16 个,另有 2 个所有 token 都经过的 shared expert。若每个 routed expert 都在 7168 维 hidden space 里做完整 FFN,专家参数和通信都会非常昂贵。

K3 让 routed 分支先把输入从 7168 维压到 3584 维 latent space,在较窄的空间里完成专家计算,再投影回 7168 维。简化后的路径是:

x [7168]
  ├─> shared expert 1/2(保持全宽)
  └─> W_down: 7168 -> 3584
        └─> 选中的 16 个 routed experts
              └─> W_up: 3584 -> 7168
                     └─> 与 shared 分支相加

为什么 shared expert 很重要?如果所有专家都在半宽 latent space,模型的每条信息通路都要经过压缩;shared expert 保留了全宽的基础通路,routed experts 则负责在 latent space 里提供稀疏的特化能力。这样既控制了 896 个专家的参数量,又没有把所有信息都押在低维瓶颈上。

5. 激活为什么会爆:先看 routed 分支的乘法链

一个 routed expert 的 FFN 大致包含 down projection、非线性、up projection,再乘上 router 权重。多个 expert 聚合后还要与 shared 分支相加。2.8T 规模下,任何一段输出的幅值失控,都可能在后面的矩阵乘中被放大。

K3 的第一道保险是在 routed experts 聚合之后、up projection 之前加 RMSNorm:

\[u=\sum_{i\in\mathrm{TopK}}p_iE_i(W_{down}x), \qquad y=W_{up}\operatorname{RMSNorm}(u)+y_{shared}\]

它不改变 token 选择哪个专家,只把 16 个专家聚合后的 scale 拉回可控范围,再交给升维矩阵。

第二道保险是把 SwiGLU 换成 SiTU-GLU。SwiGLU 的两个分支在大输入上没有硬上限,低精度训练时更容易产生极端值。SiTU-GLU 给两个分支加 soft cap:

\[f(x)=\left[\beta_1\tanh(W_gx/\beta_1)\odot\sigma(W_gx)\right] \odot\left[\beta_2\tanh(W_ux/\beta_2)\right]\]

K3 取 $\beta_1=4$、$\beta_2=25$。在输入接近 0 时,tanh 近似线性,函数仍像 SwiGLU 一样工作;输入很大时,每个分支会平滑封顶,输出幅值上界约为 $\beta_1\beta_2=100$。

SiTU-GLU 的幅值控制

看图时重点看两段:原点附近红色 SiTU-GLU 跟 SwiGLU 接近,说明普通输入区域的表达能力没有被明显改写;大输入区域红线逐渐压平,说明低精度路径不会任由异常激活继续放大。

6. 896 个专家为什么会负载不均

即使激活稳定,router 还有另一个问题:它可能反复把 token 送到少数热门专家。假设一个 batch 有 8 个 token、4 个专家,每个 token 只选 1 个,朴素 Top-k 可能得到负载:

expert 1: 4 tokens
expert 2: 3 tokens
expert 3: 1 token
expert 4: 0 token

在专家并行(EP)中,专家 1 所在的 GPU 要处理 4 个 token,专家 4 所在的 GPU 没活干;整个 batch 的时间由最忙的 GPU 决定。专家越来越多时,少量偏差更容易把某些专家推成“过热”或“濒死”。

过去常见的 auxiliary loss 会在训练目标里惩罚不均衡,或者像 DeepSeek 的 auxiliary-loss-free 方法一样维护每个专家的 bias,用正负小步更新来调节选择。K3 关注的是另一个事实:一次 router 前向已经告诉我们,每个 token 要让某个专家进入 Top-k,至少需要跨过什么门槛。

7. Quantile Balancing:从门槛直接反推 bias

对 token $i$,取它的第 $k+1$ 大 router 分数作为门槛 $\alpha_i$。专家 $j$ 如果加上 bias $b_j$,能够进入该 token 的 Top-k,当且仅当:

\[s_{i,j}+b_j>\alpha_i\]

移项以后:

\[b_j>\alpha_i-s_{i,j}\]

右侧就是“专家 $j$ 还差多少分才能进入 token $i$ 的 Top-k”。把一个 batch 中所有 token 的这个 margin 收集起来,如果希望专家 $j$ 被约 $q=mk/n$ 个 token 选中,就应该选择一个分位点作为 $b_j$。

# s: [m, n],m 个 token 对 n 个专家的 router 分数
# alpha: [m],每个 token 的 Top-(k+1) 门槛
margin = alpha[:, None] - s
b_hat = -torch.quantile(margin, q=1 - k / n, dim=0)
b = b_hat - b_hat.mean()

为什么最后要减均值?给所有专家 bias 同时加上同一个常数,不会改变 Top-k 排序;减掉均值只是去掉这个无意义的公共偏移。

以 8 token、4 expert、$k=1$ 为例,目标负载是 $q=8\times1/4=2$。Quantile Balancing 不是让 bias 每轮只移动一点,而是根据当前 margin 分布直接估计一个能让每个专家接近 2 个 token 的 bias。

Quantile Balancing 示意

图中左侧是热门专家和“死亡专家”的负载差异,右侧是加入 bias 后各专家 margin 被重新对齐。实际训练中,K3 不会把百万级 margin 全部 gather 到一张卡上,而是每张卡先统计每个专家的 1000-bin 直方图,再 all-reduce 整数计数,最后从合并直方图近似读取分位点。通信量从所有浮点 margin 降到每个专家固定数量的 bin,误差约为 bin 宽度级别。

8. 为什么 K3 的宽度设计要和稳定性绑在一起

把 routed experts 放进 0.5× latent space,并不是单纯为了少算一半矩阵。它同时降低了:

  • 896 个专家各自保存 FFN 权重的成本;
  • token dispatch 到专家后的 GEMM 规模;
  • 多专家聚合后传回全宽 hidden space 的压力。

但压缩宽度以后,shared expert、RMSNorm 和 SiTU-GLU 就变成了必要的稳定性配套;专家数量变多以后,Quantile Balancing 又承担了把计算量分回各个 rank 的职责。latent space、激活 soft cap、路由 bias 不是三个互不相关的小技巧,而是同一个“超大稀疏 FFN 如何可训练、可并行”的问题的不同答案。

论文的 scaling law 图报告 K3 相比 K2 约 2.5× 的 scaling efficiency 提升。这里的 scaling efficiency 是在相同训练损失/能力目标下所需计算量的拟合比较,不等同于单个 kernel 的 2.5× 加速;它还包含架构、训练配方和系统效率的共同影响。

K2 与 K3 的 scaling law

9. 小结

本篇可以这样记:

  1. 标准残差只把上一层送给下一层,AttnRes 把“从前面哪些层取信息”变成了可学习的 depth attention;
  2. Full AttnRes 会保存所有层,Block AttnRes 用约 8 个 block 表示保留大部分效果;
  3. 896 个 routed experts 如果都在全宽空间运行,参数和通信太贵,因此先进入 0.5× latent space;
  4. shared expert 保住全宽通路,RMSNorm 和 SiTU-GLU 控制激活幅值;
  5. Quantile Balancing 从第 $k+1$ 个分数提供的门槛反推 bias,让负载均衡从“小步试错”变成“按 margin 分布直接估计”。

下一篇进入系统部分:MoonEP 如何让专家并行的 token 数固定,KDA 的状态如何跨卡传递,1M context 的 Agentic RL 和 prefix cache 又为什么需要重新设计。


扯一句题外话:Stable LatentMoE 和量化、剪枝属于相邻但不同的问题。《动手学 AutoML:从 NAS 到大语言模型优化实战》第 8、11 章整理了 LLM 压缩和后训练剪枝,适合从模型压缩的角度补充“如何让大模型在更低成本下运行”。

动手学AutoML书籍封面

Flag Counter