专栏Kimi K3 模型结构·宽度6 / 8
11 min学习

Kimi K3 模型结构(5):宽度维度,Stable LatentMoE

896 选 16 的 MoE 为什么要在一半宽度的 latent 空间里跑。先复述 NVIDIA LatentMoE 的论证链,再逐个拆 K3 加的三样东西:latent 上的 RMSNorm、SiTU-GLU 激活、Quantile Balancing 负载均衡,包括它从平衡指派问题推出来的过程。

目录11 节

对应论文 §2.3 和附录 B、C、D,公式 (11) 到 (14)。K3 的 93 层里有 92 层的 FFN 是这个模块,它占了模型 98% 的参数。上游是 NVIDIA 2026 年 1 月的 LatentMoE,K3 在它上面加了三个东西,合起来叫 "Stable"。

文本 token图像 / 视频视觉 token 与文本 token 交错后进入同一条主干残差流(prefix sum)重复 23 次3 KDA : 1 Gated MLA第 2 – 92 层每 12 层是一个 AttnRes 块AttnRes 来源(最多 9 个)每个 α 前都有这一组来源α = softmax(wₗ · RMSNorm(来源))wₗ 是每个子层各一个的可学习向量logits输出前再聚合一次全部块最终隐藏状态 + 下一个 token 的 embeddingAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和αAttention Residuals 算子:用本子层的伪 query wₗ 对来源打分,softmax 加权求和α词表 163840,隐藏维 7168Token Embedding163840 × 7168从零训练的视觉编码器,hidden 1024,12 头MoonViT-V227 层 · 0.4B · patch 14merge_type = sd2_tpool2×2 Pixel-shuffle+ 时间池化 · token ÷4PatchMergerMLPV2,GELU,RMSNormProjector4096 → 4096 → 7168词嵌入永远是一个来源b₀ = Embedding每个块的输出是块内所有子层输出之和b₁ … bₙ₋₁已完成的块(每块 12 层之和)本块里已经算完的子层之和,相当于块内的普通残差流bₙ⁽ⁱ⁾ 当前块 partial sumKimi Delta Attention:96 头 × 128 维,状态 128×128/头KDA第 1 层第一层不用 MoE,用一个稠密 SiTU-GLU FFNDense FFN仅第 1 层 · 中间维 33792三层 KDA,每层后接一个 Stable LatentMoEKDA×37168 → 3584 latent → 16 个 routed expert → RMSNorm → 7168Stable LatentMoE896 选 16 + 2 sharedDeepSeek 式 MLA,无位置编码,满秩 sigmoid 输出门Gated MLA×1 · NoPE同上Stable LatentMoE896 选 16 + 2 shared主干末尾额外放一层全局注意力Gated MLA第 93 层 · 收尾同上Stable LatentMoE第 93 层最终归一化RMSNorm不与 embedding 共享权重LM Head→ 163840预训练时 1 层多 token 预测;post-training 微调成投机解码的 draftMTP 层 ×1镜像主干 block · 部署时做 EAGLE-3 draft
序列 · KDA序列 · Gated MLA深度 · AttnRes宽度 · Stable LatentMoE输入 · MoonViT-V2零件
你现在在这里:每层注意力后面的 FFN。第 1 层是稠密 FFN,其余 92 层都是 Stable LatentMoE。

数字

Kimi K2Kimi K3
routed expert 数 NN384896
每 token 激活 kk816
shared expert12
稀疏度 N/kN/k4856
expert 的输入宽度71683584(latent)
expert 中间维20483072
激活函数SwiGLUSiTU-GLU

expert 数翻了一倍多,每个 token 用的 expert 也翻倍,但每个 expert 只看一半宽度的输入。这一组数字是 LatentMoE 论文推荐的配方。

LatentMoE 的论证链

LatentMoE 不是一篇扫超参的论文,是一篇系统和算法共同设计的论文。它的问题是:MoE 层的几个旋钮里(expert 数 NN、激活数 kk、宽度 dd、expert 中间维 mm、shared expert 数),哪个该减、哪个该加。论证分五步:

  1. 推理时 expert 处在访存瓶颈。 每个 expert 一步只处理几百个 token,算术强度远低于 GPU 的拐点,时间花在读权重上。所以要优化的是「每个参数的准确率」,不是「每个 FLOP 的准确率」。
  2. 通信量正比于 kdk \cdot d expert 并行的 all-to-all 每个 token 要发 kkdd 维向量。mm 帮不上忙,能砍的只有 ddkk
  3. 非线性预算正比于 kmk \cdot m 单隐层网络的逼近误差按 1/u1/u 下降,uu 是非线性单元数,和输入维数无关。所以 kkmm 不能减。
  4. 特征秩是宽度的下限。 任务有内在的特征秩 reffr_{\text{eff}},宽度低于它就崩。实验里 reffd/4r_{\text{eff}} \le d/4
  5. 组合稀疏性。 (αNαk)(Nk)α\binom{\alpha N}{\alpha k} \ge \binom{N}{k}^\alpha,把 NNkk 同时放大 α\alpha 倍,expert 组合数超指数增长。

综合:只砍 dd,把 routed expert 的输入宽度从 dd 降到 =d/α\ell = d/\alpha,省下来的通信和权重带宽换成 α\alpha 倍的 expert 数和 α\alpha 倍的 top-kk。推理成本回到原点,准确率靠 4、5 两条涨。原论文推荐 α=4\alpha = 4,验证损失在 α4\alpha \le 4 时持平;K3 用了保守的 α=2\alpha = 2

原论文的实验:95B 模型在 300B token 上,MMLU-Pro 从 29.3 涨到 34.9,激活参数不变。

K3 的前向

论文公式 (11):

u=iTk(x)piEirouted(Wx),y=j=1NsEjshared(x)+WRMSNorm(u).u = \sum_{i \in \mathcal{T}_k(x)} p_i\, E_i^{\text{routed}}(W^{\downarrow} x), \qquad y = \sum_{j=1}^{N_s} E_j^{\text{shared}}(x) + W^{\uparrow}\operatorname{RMSNorm}(u).
x7168W↓7168→3584中间3072Σ pᵢ Eᵢ(z)3584RMSNormK3 新增W↑3584→7168y716816 个 routed expert(从 896 个里选)每个:3584 → 3072 → 3584,SiTU-GLU在 latent 空间里 dispatch、计算、加权求和z = W↓x一层只有一个一层只有一个router 看的是全宽 x:sigmoid(W_r x),896 个分数,加偏置选 162 个 shared expert全宽 7168 → 6144 → 7168相加
矩形高度按维度画。routed 路在 3584 维的 latent 空间里做所有事:分发、16 个 expert 的计算、加权求和,然后再展开回 7168。绿色 RMSNorm 是 K3 加的。shared expert 走全宽,直接加到输出上。

对着 HF 代码走一遍:

python
def forward(self, hidden_states):                    # KimiSparseMoeBlock
    identity = hidden_states
    topk_idx, topk_weight = self.gate(hidden_states)  # router 看全宽的 x,7168 → 896 分
    hidden_states = self.routed_expert_down_proj(hidden_states)   # 7168 → 3584,一层一个
    y = self.moe_infer(hidden_states, topk_idx, topk_weight)      # 16 个 expert 在 3584 维上算、加权求和
    y = self.routed_expert_norm(y)                    # RMSNorm(3584)      ← K3 新增
    y = self.routed_expert_up_proj(y)                 # 3584 → 7168,一层一个
    return y + self.shared_experts(identity)          # 2 个 shared expert 走全宽,直接加
  • router 在全宽的 xx 上打分,不在 latent 上。这是 LatentMoE 的设计,K3 保留。
  • WW^{\downarrow}WW^{\uparrow} 每层只有一个,所有 routed expert 共用。这和 MoLAE 那种每个 expert 各自低秩分解不同,只有共享的投影才能让 all-to-all 也变窄。
  • 加权求和发生在 latent 空间WW^{\uparrow} 每 token 只作用一次,在 all-to-all combine 之后。所以 dispatch 和 combine 两趟通信都是 3584 维。
  • 每个 routed expert 是 3584 → 3072 → 3584 的三矩阵 GLU(KimiBlockSparseMLP,gate、up、down)。
  • 两个 shared expert 在代码里合成了一个 KimiMLP,中间维 3072 × 2 = 6144,全宽 7168 → 6144 → 7168。数学上等价于两个并联的 3072 维 expert 相加。

为什么需要 "Stable"

原论文的 routed 路径上没有任何归一化WWdown(i)σ(Wgate/up(i)Wx)W^{\uparrow} \cdot W_{\text{down}}^{(i)} \cdot \sigma(W_{\text{gate/up}}^{(i)} W^{\downarrow} x),四个矩阵乘串起来中间只有一个激活函数。WW^{\downarrow}WW^{\uparrow} 同时从 16 条路径拿梯度。原论文没有讨论这条链的条件数、初始化或者中间的归一化。

K3 论文说,在 56 的稀疏度和 2.8T 的规模下,这个设计的两个问题被放大了:routed 分支内部的激活爆炸;接近 10310^3 个 expert 的负载均衡超出了现有无辅助损失偏置更新的适用范围。三个对策各自对应:RMSNorm 和 SiTU-GLU 管激活,Quantile Balancing 管均衡。

Stable 之一:latent 上的 RMSNorm

原 LatentMoE 直接把 WW^{\uparrow} 作用在聚合结果 uu 上,而 uu 的尺度随着选中了哪些 expert、路由权重多大而变。K3 在聚合和上投影之间插一个 RMSNorm(latent_moe_use_norm = true,代码里 routed_expert_norm,3584 维)。作用是让 routed 分支在和全宽的 shared 分支相加之前,对尺度变化不敏感。论文说除了稳定训练,这个 RMSNorm 还一致地改善了验证损失和下游指标。

Stable 之二:SiTU-GLU

先看 SwiGLU:Swish(Wgx)Wux=(Wgx)σ(Wgx)Wux\operatorname{Swish}(W_g x) \odot W_u x = (W_g x)\,\sigma(W_g x) \odot W_u x。两个乘数都无界。当两边同时出现大值时,乘积产生激活离群点,低精度下容易溢出。原始 GLU 用 sigmoid 门避免了门的无界增长,但失去了 Swish 在正半轴近似线性的响应。

SiTU-GLU(论文公式 (12))对 Swish 门里的线性因子和上支分别做平滑截断 softcap(x,β)=βtanh(x/β)\operatorname{softcap}(x, \beta) = \beta\tanh(x/\beta)

SiTU-GLU(x)=[β1tanh ⁣(Wgxβ1)Sigmoid(Wgx)][β2tanh ⁣(Wuxβ2)],β1=4, β2=25.\operatorname{SiTU\text{-}GLU}(x) = \Big[\beta_1 \tanh\!\Big(\frac{W_g x}{\beta_1}\Big) \odot \operatorname{Sigmoid}(W_g x)\Big] \odot \Big[\beta_2 \tanh\!\Big(\frac{W_u x}{\beta_2}\Big)\Big], \qquad \beta_1 = 4,\ \beta_2 = 25 .
x(同时喂给门和上支)f(x)
GLU:σ(x)·xSwiGLU:x·σ(x)·xSiTU-GLU:β₁tanh(x/β₁)·σ(x) · β₂tanh(x/β₂)上界 β₁β₂ = 100

附录 B 给的三条性质:

  • 一阶等价 SwiGLU。 βtanh(z/β)=z+O(z3/β2)\beta\tanh(z/\beta) = z + O(z^3/\beta^2),原点附近和 SwiGLU 一样;β1,β2\beta_1, \beta_2 \to \infty 时逐点收敛到 SwiGLU。
  • 输出有界。 tanh<1|\tanh| < 10<σ<10 < \sigma < 1,所以每个输出坐标的绝对值不超过 β1β2=100\beta_1 \beta_2 = 100
  • 比硬截断好。 平滑的 cap 在远离饱和边界的地方保持非零梯度,硬 clamp 做不到。

sigmoid 因子保留着,所以负半轴的响应仍然趋于零,Swish 的形状没丢。config 里 hidden_act = situactivation_situ_beta = 4.0activation_situ_linear_beta = 25.0。注意这个激活不只用在 routed expert 上:代码里 KimiMLP 也按 hidden_act 选激活,所以第 1 层的稠密 FFN 和 shared expert 用的也是 SiTU-GLU。

python
class SituAndMul(nn.Module):
    def forward(self, x):
        d = x.shape[-1] // 2
        gate, up = x[..., :d].float(), x[..., d:].float()
        g = self.beta * torch.tanh(gate / self.beta) * torch.sigmoid(gate)   # β₁ = 4
        u = self.linear_beta * torch.tanh(up / self.linear_beta)             # β₂ = 25
        return (g * u).to(x.dtype)

路由:sigmoid 打分,加偏置选,按原分归一化

论文公式 (13):

si=Sigmoid(Wrxi),Ti=argtopk(si+b),pi,j=si,jrTisi,r, jTi.s_i = \operatorname{Sigmoid}(W_r x_i), \qquad \mathcal{T}_i = \operatorname{argtop}_k(s_i + b), \qquad p_{i,j} = \frac{s_{i,j}}{\sum_{r \in \mathcal{T}_i} s_{i,r}},\ j \in \mathcal{T}_i .

三点:打分是 sigmoid 不是 softmax(moe_router_activation_func = sigmoid);选哪些 expert 看加了偏置的分数,但混合权重 pp 用的是原始分数,所以偏置只管分发、不改变 expert 输出的加权,也不进 router 的梯度;选中的 16 个权重重新归一化到和为 1(moe_renormalize = true)。这就是 DeepSeek-V3 的 aux-loss-free 路由框架,config 里 topk_method = noaux_tc,那个偏置向量就是 e_score_correction_bias,推理时冻结。K3 的 num_expert_group = 1,没有分组限制路由。

Stable 之三:Quantile Balancing

原方法的问题。 DeepSeek-V3 用固定步长更新偏置:bjbj+γsign(ˉj)b_j \leftarrow b_j + \gamma\,\operatorname{sign}(\bar\ell - \ell_j),负载高于平均就减、低于就加。γ\gamma 大了震荡,小了追不上。expert 到 896 个之后,这个折衷变得难调;不平衡的路由拖慢 expert 并行训练,还可能让一些 expert 训练不足。

QB 的做法:直接算出让每个 expert 恰好拿到目标负载的偏置。 设一个 batch 有 mm 个 token、nn 个 expert、每 token 选 kk 个,目标负载 q=mk/nq = mk/n。一次前向里:

  1. 路由时不选 top-kk 而选 top-(k+1)(k+1)。前 kk 个是真正走的路,第 k+1k+1 个的分数记作 αi\alpha_i,它是一个 expert 想进 token ii 的 top-kk 必须超过的截止线
  2. 固定所有截止线,问:expert jj 的偏置取多少,恰好有 qq 个 token 会选它?token ii 选 expert jj 当且仅当 si,j+b^j>αis_{i,j} + \hat b_j > \alpha_i,即 margin si,jαis_{i,j} - \alpha_i 超过 b^j-\hat b_j。让恰好 qq 个 margin 超过阈值,阈值就是第 (q+1)(q+1) 大的 margin,也就是 margin 的 (1k/n)(1 - k/n) 分位数:
b^j(t+1)quantile1k/n(s:,jα(t)),b(t+1)b^(t+1)mean(b^(t+1))1.\hat b_j^{(t+1)} \leftarrow -\operatorname{quantile}_{1-k/n}\big(s_{:,j} - \alpha^{(t)}\big), \qquad b^{(t+1)} \leftarrow \hat b^{(t+1)} - \operatorname{mean}(\hat b^{(t+1)})\,\mathbf{1}.

第二行减掉公共偏移,因为对 top-kk 选择没影响。新偏置下一步才生效,一个 batch 永远不会用从它自己算出来的偏置路由。

(a) 直接 Top-1:负载 4 / 3 / 1 / 0t1t2t3t4t5t6t7t8E14E23E31E40(c) 加 QB 偏置后:负载 2 / 2 / 2 / 2t1t2t3t4t5t6t7t8E12b = -0.19E22b = +0.03E32b = +0.04E42b = +0.13
m = 8 个 token,n = 4 个 expert,k = 1,目标负载 q = mk/n = 2。左边按原始分数路由,E1 过热、E4 一个都没有。右边每个 expert 的偏置取自本列 margin 的第 (q+1) 大值,一步就到 2/2/2/2。红色是被 QB 改动的边。分数是示意用的手工数据。

为什么这是「对」的答案:附录 C 的推导。 从最大分数的平衡指派问题出发:

maxxi,j{0,1}i,jxi,jsi,js.t.jxi,j=k,ixi,j=mkn.\max_{x_{i,j} \in \{0,1\}} \sum_{i,j} x_{i,j}\, s_{i,j} \quad\text{s.t.}\quad \sum_j x_{i,j} = k,\qquad \sum_i x_{i,j} = \frac{mk}{n}.

放松成线性规划(二分 bb-匹配多面体是整的,放松无损)。对两组等式约束引入乘子 αi\alpha_i(token 侧)和 βj\beta_j(expert 侧),交换 min 和 max,内层对每个 xi,jx_{i,j} 独立:si,jαiβjs_{i,j} - \alpha_i - \beta_j 为正就取 1,否则取 0。代回去得到凸的对偶目标

minα,β i,jmax(0, si,jαiβj)+kiαi+mknjβj.\min_{\alpha, \beta}\ \sum_{i,j}\max(0,\ s_{i,j} - \alpha_i - \beta_j) + k\sum_i \alpha_i + \frac{mk}{n}\sum_j \beta_j .

对它做坐标下降。固定 β\betaαi\alpha_i:目标对 α\alpha 分段线性,斜率是 kk 减去「超过 α\alpha 的 margin 个数」,所以恰好 kk 个 margin 在 α\alpha 之上时取最小,闭式解是 siβs_i - \beta 的第 (k+1)(k+1) 大值。对称地,固定 α\alphaβj\beta_j,闭式解是 s:,jαs_{:,j} - \alpha 的第 (mk/n+1)(mk/n + 1) 大值。两边都是同一个 (1k/n)(1-k/n) 分位数,方法因此得名。

最优解处 xi,j=1x^*_{i,j} = 1 当且仅当 si,jαiβj>0s_{i,j} - \alpha_i^* - \beta_j^* > 0,结合 token 侧约束,选中的正好是 siβs_i - \beta^* 的 top-kk所以路由只需要 expert 侧的 β\betab=βb = -\beta),token 侧的 α\alpha 是随 batch 变的中间量,用完就丢。 部署时是固定 top-kk 加冻结偏置,不算任何分位数。

和 DeepSeek-V3 的关系。 expert 侧子问题的次梯度是 mkni1[si,jαiβj>0]\frac{mk}{n} - \sum_i \mathbf 1[s_{i,j} - \alpha_i - \beta_j > 0],即目标负载减实际负载。对这个目标做 SignSGD,就是 DeepSeek-V3 的 sign 更新(差一个 b=βb = -\beta 的符号约定)。sign 更新只保留了负载误差的方向,QB 直接跳到同一个对偶目标的精确坐标极小点。这解释了 QB 为什么没有步长这个超参,以及为什么接近 10310^3 个 expert 也能在几步内平衡。

直方图估计(附录 D)。 分位数要在整个全局 batch 上取,margin 有几百万个,散在各 rank 和梯度累积步里,收齐再排序不现实。观察:更新只需要每个 expert 的 margin 分布,不需要值本身。做法:

  • 对「所需偏置」ri,j:=αisi,jr_{i,j} := \alpha_i - s_{i,j} 做直方图。它的范围有界:s(0,1)s \in (0,1)αi\alpha_i 是某个 expert 加偏置后的分数,所以 r[bmin1, bmax+1]r \in [b_{\min} - 1,\ b_{\max} + 1],每步按当前偏置重算范围,分 BB 个均匀 bin。
  • 前向时每个 rank 把本地 ri,jr_{i,j} scatter-add 进 n×Bn \times B 的计数矩阵,跨 micro-batch 累加,不通信。步末一次整数 all-reduce 求和,所有 rank 从同一份全局直方图读分位数:找累计计数首次达到 q\lceil q \rceil 的 bin,在 bin 内线性插值。
  • 误差不超过 bin 宽,B=1000B = 1000 时是几个 10310^{-3};通信是每层每步 nBnB 个整数,和 mm 无关,比每个 micro-batch 交换原始 margin 便宜两个数量级;计数可加,所以结果是全局 batch 的分位数而不是各 rank 分位数的平均。

中文读者可以先看苏剑林的《MoE 环游记》第 6 篇(spaces.ac.cn/archives/11619),QB 的思路那里有更慢的铺垫。

参数账

这一层占了 K3 的绝大部分参数,算一下:

部件每层参数层数合计
routed expert:3 × 3584 × 3072 × 89629.6B922.72T
shared expert:3 × 7168 × 6144132M9212.2B
WW^{\downarrow} + WW^{\uparrow}:2 × 7168 × 358451.4M924.7B
router:896 × 71686.4M920.6B

routed expert 一项就是 2.72T,占 2.78T 的 98%。每 token 激活的部分:16 个 expert 528M,加 shared 132M,加投影 51M,每层约 711M,92 层约 65B;剩下的 39B 激活参数在注意力(KDA 69 层约 30B,MLA 24 层约 6B)、第 1 层稠密 FFN(0.7B)、embedding 和 LM head(各 1.2B)里,合计和论文的 104.2B 对得上。

部署时只有 routed expert 的权重量化到 MXFP4(config 的量化配置里排除了注意力、shared expert、latent 投影、LM head 和视觉部分),激活用 MXFP8,整个 post-training 阶段做量化感知训练。

术语坑

  • LatentMoE 原论文的 95B 模型 expert 是无门的 Squared-ReLU 两矩阵,K3 是带门的三矩阵。比参数量时别混。
  • 原论文推荐 α=4\alpha = 4,K3 用 2。按原论文的配方,K3 相当于把一个 448 选 8、宽 7168 的 MoE 压半宽后翻倍。
  • e_score_correction_bias 就是 QB 的偏置,训练完冻结。它不参与 pi,jp_{i,j} 的计算。
  • shared expert 在代码里是一个 6144 维的 MLP,论文写的是 Ns=2N_s = 2 个。两者等价。

下一篇

三个维度讲完了。第 6 篇收零件:输入端的 MoonViT-V2 和投影层,MTP 层怎么变成 EAGLE-3 的 draft,以及 Per-Head Muon。