一、一个被绑死的耦合关系
第 10 讲讲的前馈层,FFN(x) = σ(xW₁+b₁)W₂+b₂,是一个**稠密(dense)**的计算:不管输入的 token 是什么,W₁、W₂ 里的每一个参数都会参与这次计算。
这带来一个直接的后果:想让模型"知道得更多"(参数更多),就必须让每一个 token 的计算成本跟着等比例上涨——这两件事在标准的稠密 Transformer 里被死死绑在一起,没有办法只增加"知识储备"而不增加"每次调用的开销"。
但这个耦合真的是必须的吗?看一个类比就知道未必。
二、打个比方:大医院的会诊制度
一家大型综合医院可能有上百位不同科室的专科医生——心内科、骨科、皮肤科、神经科……医院整体的诊疗能力(总专家数)非常庞大。
但一个因为崴了脚来就诊的病人,不需要走遍全部一百个科室,他只需要挂骨科(也许再加一个放射科拍片子),其余九十几位专家这次完全没有介入他的诊疗——不消耗他们的时间,也不产生任何费用。
⭐ 医院整体的"知识储备"可以做得很大,但每个具体病人实际消耗的诊疗资源,只取决于他这次真正用到的那几位专家,和医院总共请了多少专家没有直接关系。这正是这一讲要讲的机制——混合专家(Mixture of Experts,MoE)——想要解开的那个耦合关系。
三、把直觉钉成机制
从一个前馈层,变成一组并行的"专家"
MoE 的做法:把第 10 讲的一个前馈层,换成 N 个结构相同、但参数各自独立的前馈层,叫专家(expert),记作 E₁, E₂, ..., E_N。再加一个额外的、小巧的组件——路由器(router,也叫 gating network),它的任务是:看一眼这个 token 的向量,判断该把它交给哪几个专家处理。
路由器本身的结构,就是第 5 讲讲过的"一层"——一次矩阵乘法,把输入向量投影成 N 个数字(每个数字对应一个专家的"匹配分数"),再过一次 softmax(第 5 讲)变成合法的权重:
gate = softmax(x · W_gate) ← 每个专家一个分数,N 个数字加起来等于 1
关键的一步在这里:不是所有专家都会参与这次计算。路由器只挑出分数最高的 k 个专家(k 远小于 N,比如 N=8 个专家里只挑 k=2 个),只有这 k 个专家真正跑一遍前馈层的计算,其余 N-k 个专家这一次完全不参与,不消耗任何计算。选中的专家各自的输出,按(重新归一化过的)路由权重加权求和,作为这一层最终的输出:
Output(x) = Σᵢ∈TopK gate'ᵢ × Eᵢ(x)
混合专家(MoE):把一个前馈层替换成多个参数独立的专家前馈层,加一个路由器为每个 token 挑选其中一小部分专家参与计算,模型的总参数分布在全部专家里,但每次前向传播实际使用的参数只是其中一小部分。
一个可以手算的路由过程
用一个简化到可以完全手算的例子走一遍。设有 N=3 个专家,每个专家为了让手算干净,简化成一个把 2 维输入变成一个标量输出的线性函数:
x = [1, 0]
E₁(x) = x·[2,0]ᵀ = 2×1+0×0 = 2
E₂(x) = x·[0,3]ᵀ = 0×1+3×0 = 0
E₃(x) = x·[1,1]ᵀ = 1×1+1×0 = 1
路由器算出的打分(直接给出投影后的结果,矩阵乘法本身第 3 讲已经讲过,这里重点在路由逻辑):
gate_logits = [5, 1, 3] ← 专家 1、2、3 各自的原始打分
softmax([5,1,3]):
exp(5)≈148.4,exp(1)≈2.72,exp(3)≈20.09,总和≈171.2
gate₁ ≈ 148.4/171.2 ≈ 0.867
gate₂ ≈ 2.72/171.2 ≈ 0.016
gate₃ ≈ 20.09/171.2 ≈ 0.117
取 k=2(top-2):分数最高的两个专家是专家 1(0.867)和专家 3(0.117),专家 2(0.016)被丢弃,完全不参与这次计算。把选中的两个专家的权重重新归一化(让它们自己加起来等于 1):
gate₁' = 0.867/(0.867+0.117) ≈ 0.881
gate₃' = 0.117/(0.867+0.117) ≈ 0.119
最终输出 = 0.881×E₁(x) + 0.119×E₃(x) = 0.881×2 + 0.119×1 = 1.881
⭐ 专家 2 在这次计算里,从头到尾没有被调用过一次——它的参数原封不动地待在那里,这次的前向传播完全没有触碰它们,自然也不消耗对应的计算量。如果把专家 2 换成一个参数量更大、能力更强的专家,只要它没被这次路由选中,这次调用的计算成本一分不会增加——这正是"参数量和计算量解耦"的具体样子:模型的总参数量(三个专家的参数总和)可以做得很大,但每个 token 实际消耗的计算量,只取决于被选中的 k 个专家。
四、这一讲机制自带的代价
⭐ 参数量和计算量解耦了,但显存/存储量没有解耦:虽然每个 token 只激活 k 个专家,但没有人能提前知道下一个 token 会被路由到哪几个专家——为了让任何一个 token 随时都能被正确路由,全部 N 个专家的参数必须始终完整地待在显存里,不能因为"这次用不到"就先挪走。第四节医院的比方在这里要打一个补丁:医院可以不给没被挂号的科室发工资(省计算),但没法把没被挂号科室的医生和设备直接搬空(省不了场地/显存)。
⚠️ 路由决策是离散的,给训练带来了真实的困难:第 6 讲讲过,梯度下降靠的是连续的、处处可求导的函数。“选出分数最高的 k 个专家"这个 top-k 操作本身是不连续的——分数差一点点,选中的专家集合可能完全不同,这种跳变没法用普通的求导规则处理。实践中要靠一些近似和平滑技巧绕开这个障碍,这也是 MoE 训练比稠密模型更容易不稳定的原因之一。
⚠️ 负载不均衡是一个真实存在、需要专门对付的工程问题:训练初期,如果某个专家因为参数随机初始化的偶然性,恰好对某一类 token 处理得稍微好一点,路由器会倾向于把更多这类 token 继续分给它——这个专家因此获得更多训练、变得更强,进一步吸引更多路由,形成"赢家通吃"的正反馈循环,最终可能变成少数专家超载、其余专家几乎无人问津,总参数量虽然很大,但实际利用率很低。实践中通常需要在训练目标里额外加一项"负载均衡损失”,强制路由器把 token 更均匀地分给所有专家,这不是一个理论上完美、一劳永逸的方案,是一个需要持续调校的工程约束。
⚠️ 专家的分工不保证是人类能读懂的"专业化":不能想当然地假设"专家 1 负责语法、专家 2 负责事实"这种整齐的分工——真实训练出来的专家分工往往是统计意义上的、不直观的模式,第 31 讲讨论可解释性边界时,这是同一类"能用,但说不清楚具体在干什么"的例子。
五、和后面课程的关系
- 第 14、15 讲讲预训练数据规模和 Scaling Laws 时,MoE 会重新出现——它是"想要更大的模型、但计算预算有限"这个约束下,一种具体的参数效率手段,服务于第 1 讲就提出的"规模换效果"这个赌注。
- 第 22 讲讲推理阶段的显存优化时,这一讲第四节"计算量解耦了、显存没解耦"的代价会重新变成一个具体的工程账——推理服务必须为全部专家预留显存,即便每个请求只用得上其中一小部分。
- 第 31 讲讨论可解释性时,“专家分工不直观"是同一类问题在这门课里反复出现的又一个例子。
六、本讲小结
- 稠密前馈层把参数量和计算量绑死了:每个 token 都要激活全部参数,想要模型"知道更多"就必须让每次调用的开销等比例上涨。
- ⭐ MoE 把一个前馈层换成多个并行的专家,加一个路由器(本质是第 5 讲的"一层+softmax”)给每个 token 挑出分数最高的
k个专家参与计算,其余专家这次完全不参与——本讲手算的例子具体验证了打分、top-k 选择、重新归一化、加权求和的完整流程,以及"没被选中的专家全程零参与"这个核心特性。 - ⭐ “参数量和计算量解耦” 说的是:模型的总参数(全部专家之和)可以很大,但每个 token 实际付出的计算只取决于被选中的少数专家。
- ⚠️ 但显存/存储量没有被解耦——全部专家必须随时待命在显存里,因为路由结果无法提前预知。
- ⚠️ 训练上的真实代价:top-k 选择不连续、给反向传播带来困难;负载不均衡容易导致"赢家通吃",需要额外的均衡损失来强制干预;专家的分工不保证符合人类直觉的可解释性。
思考题
- 本讲第三节的例子里,如果把
k从 2 改成 1(只选分数最高的专家),最终输出会变成什么?和k=2时的1.881相比,计算量和结果分别有什么变化? - 为什么说"top-k 选择是不连续的"?如果专家 1 和专家 3 的打分只差 0.0001,
k=1时选中的专家会不会因为这一点点差异发生突变?这种突变对第 6 讲讲的梯度下降有什么影响? - “负载不均衡"的正反馈循环具体是怎么发生的?请用自己的话,按"初始随机性 → 路由器的选择 → 训练信号分布 → 专家能力变化 → 又影响路由器的选择"这条链条重新描述一遍。
- 为什么说"计算量解耦了,但显存没有解耦”?如果一个 MoE 模型有 8 个专家、每次只激活 2 个,和一个参数量等于"2 个专家之和"的稠密模型相比,两者在计算量和显存占用上分别是什么关系?
- 结合第 1 讲"规模换效果"的赌注和这一讲的机制,你觉得 MoE 更适合解决"训练成本太高训不动"这个问题,还是更适合解决"推理时响应太慢"这个问题?两者是不是同一回事(提示:想一想训练和推理各自主要受什么限制)?