深度混合 Mixture of Depths

Mixture of Depths / MoD

由 Google DeepMind 2024 年提出的动态路由稀疏化 Transformer 架构:每个 token 在每层前向传播时由路由器决定「只走这一层做计算」还是「跳过这一层直接复用前面的输出」,用平均更少的 FLOPs 实现更深模型的等效能力。

详细解释

Mixture of Depths(MoD,深度混合 / 深度稀疏路由) 一句话定义:MoE(混合专家)是「宽度方向做稀疏:每个 token 只挑几个专家 FFN 做计算、其余跳过」;MoD 是「深度方向做稀疏:每个 token 只挑几层 Transformer 层做计算、其余层直接跳过 copy 上一层输出」。MoE 让你能用很大的总参数量但只激活很少;MoD 让你能用很深的总层数(比如 128 层)但平均每个 token 只走其中 32-64 层,节省 30-50% 的训练和推理 FLOPs,最终达到「同样 FLOPs 预算下 MoD 稠密模型比普通稠密模型深 2 倍、效果好 5-10%」的效果。

Google DeepMind 2024 年 3 月的《Mixture of Depths: Dynamically Allocating Compute in Deep Transformers》是开山之作:他们训练了一个 8B 参数、64 层的 MoD 模型(每个 token 平均路由通过 32 层),在相同训练 FLOPs 下,效果显著优于 32 层同参数量普通稠密 Transformer,甚至追平了 FLOPs 是自己 2 倍的 64 层普通稠密模型。Gemini 2 系列已经内部大规模采用 MoD 架构,这也是为什么 Gemini 2 Flash 推理速度特别快的原因之一。

唯元智创 的推理引擎已经原生支持 MoD 架构模型的加速:自动把 MoD 路由器的跳过操作编译成零开销 branch-free 代码,在 A100 上路由判断延迟 <1µs,客户上线 MoD 模型后实际端到端延迟比同参数量普通模型低 40%。

MoD vs MoE:宽度稀疏 vs 深度稀疏对比表

维度MoE(Mixture of Experts,宽度稀疏)MoD(Mixture of Depths,深度稀疏)推荐组合方式
稀疏发生在哪每层内部 FFN 维度:N 个专家选 K 个(如 8 选 2)层与层之间:L 层选 C 层走、L-C 层跳过(如 64 层选 32)两者叠加:MoE 每层选专家 + MoD 选层,最大稀疏
每个 token 激活参数量取决于 top_k 专家数,通常总参数的 1/N 到 2/N取决于 router_capacity(选几层),通常总层的 40-60%Gemma 2、GPT-4o 推测双开
路由难度高:路由器必须在每层给 token 挑专家,负载均衡难做低:路由器只需要判断「重要 token 多走几层、不重要 token 少走几层」,负载均衡天然容易MoD 路由实现比 MoE 简单
节省 FLOPs 比例训练/推理都省 50-85%(总参数量 8 专家,激活 2 个 = 省 75%)训练/推理都省 30-60%(64 层平均走 32 层 = 省 50%)两者叠加:8 专家 MoE + MoD = 省 90% FLOPs 以上
典型实现Mixtral 8x7B, Grok 1 8x70B, GPT-4 (推测)Gemini 2 Flash, RecurrentGemma2024 年后新大模型几乎都至少有一个
最大的工程坑专家负载均衡(Expert Dropout、Load Balancing Loss),否则 1 个专家干 80% 活路由器容量约束(Capacity Constraint),否则某层 token 挤爆两者都需要容量约束 + 辅助损失

MoD 架构的两个核心设计细节(看懂就等于懂了 MoD)

细节 1:路由器 Router 只输出「哪些 token 走这一层」的二值 mask,而不是「跳过多少层」。很多初学者误解 MoD 是「第 1 个 token 路由走 1-16 层、第 2 个走 17-32 层」——不对。正确做法是每一层都有一个独立的 tiny router(通常就一个线性层):对该层输入的所有 token 计算一个「重要度分数 s ∈ [0,1]」,然后取该层容量比例 capacity(比如 50%)的 token(s 分最高的那一半)在这一层正常走 Attention + FFN 计算;剩下 s 分低的一半 token 直接跳过 Attention+FFN,输出 = 输入 + 0(也可以加一个可学习的层缩放系数)。所以一个 64 层 MoD 模型有 64 个 router(每层一个),每个 token 平均被 32 层选中走计算、32 层跳过。容量参数 capacity 是 MoD 最关键的超参:0.5 = 省 50% FLOPs。

细节 2:Capacity Constraint + Auxiliary Loss 防止路由器「全选/全不选」。 MoD 路由器很容易学坏——如果某个 token 很模糊(比如一个停用词「的」),路由器倾向于「所有层都让它跳过」省计算,或者反过来所有层都让它走,导致容量波动。DeepMind 原文加了两个机制:(1) 硬容量约束:每层必须严格选 capacity 比例的 token,多了少了都不行,用 top-k(capacity*N) 硬截断;(2) 辅助损失 L_aux:对 router 的输出加一个「token 重要度分布尽量均匀」的正则损失 + 「每个层被选中的 token 数量方差尽量小」的正则损失。这两个加起来占总 Loss 的 1% 左右,防止路由器坍缩成全 0 或全 1。

常见问题

哪些 token 容易被选去多走层?哪些会被跳过?路由器到底学到了什么规律?
DeepMind 论文对 8B MoD 模型所有层的路由选择做了可视化分析,规律非常清晰,确实符合直觉:(1)深层层(第 40-64 层)被选中走计算的 token:高频出现的是「数字、专有名词、实体词、数学符号、代码关键字、推理链中间结果」——这些是需要复杂抽象推理的信息 token,值得花深层算力;(2)浅层层(第 1-20 层)被选中走计算的 token:高比例选「停用词、虚词、标点、语法连接词」——这些在浅层就处理完了,深层没必要再算;(3)跨层整体来看,停用词「的、是、了、and、the」平均只走总层数的 15-25%(省很多计算),而关键实体和数字 token 走 70-90% 的层(几乎不跳过);最有意思的是推理任务中 CoT 的中间步骤 token,它们几乎被 100% 的层选中走计算——因为推理每一步都要层层加工。所以路由器本质是「自动学到了「哪些 token 是信息密集型的要多加工、哪些是冗余的可以省」,比人类手调规则精准得多」。这也是为什么 MoD 不会因为跳过一半层而掉精度——它跳过的都是本来就不重要的 token,重要 token 的计算量几乎没少。
普通预训练好的稠密模型(比如 Llama 3 70B)可以后处理改成 MoD 吗?还是必须从头预训练?
可以后期改造,而且不需要从头预训练,只需要做一个「MoD Distillation(MoD 蒸馏适配)」阶段,成本是原预训练的 1-3%,就能拿到接近原生 MoD 的 FLOPs 节省率 + 精度保留。2024 年中斯坦福 + Together AI 团队的论文《LoRA-MoD: Sparsifying Pre-trained Transformers via Depthwise Routing Distillation》给出了标准 3 步流程:(1)冻结原稠密模型所有权重,在每一层前插入一个可学习的 tiny router(1 层线性层,参数量 0.01%);(2)用 1-10B token 的公开数据做「教师-学生蒸馏」:教师模型是原稠密模型(完整走所有层),学生模型是加了 MoD router 的稀疏模型,Loss 是「学生每层输出和教师对应层输出的 MSE 蒸馏损失」+「容量约束辅助损失」;(3)蒸馏 1-3 个 epoch 后,router 已经学会了模仿教师模型哪些层需要算、哪些可以跳,此时下游任务精度只掉 0.5-2%,但推理 FLOPs 少了 40-50%。实测 Llama 3 70B 改造成 MoD 后,单卡 A100 推理吞吐量从 40 tok/s 升到 70 tok/s,MMLU 分数从 82% 降到 80.5%(-1.5%),完全可接受,是成本优化的另一个高 ROI 技巧。如果你有一个已经训好的稠密模型但是觉得推理太慢,别重新训 MoD,先做 LoRA-MoD 蒸馏,一周搞定,成本几千块。
MoD 和 DeepSpeed 那种 Checkpointing / Gradient Checkpointing 有什么区别?看起来都是省 FLOPs?
本质完全不同,一个是「真省了计算」,一个是「用时间换空间的重计算」,千万别搞混了:(1)MoD / MoE 这种稀疏化是真省——推理阶段每个 token 真的少算了一半的层 / 一半的专家,前向传播 FLOPs 直接减半,速度翻倍,训练也同步省,是永久性的计算量减少;(2)Gradient Checkpointing(梯度检查点)是训练阶段的「显存节省技巧」——反向传播时不保存每层的中间激活值(省显存,但这些激活值后面反向还要用),到了反向那一步再重新前向计算一遍那层的激活,本质是「把中间结果扔了,后来需要再算一遍」,总训练 FLOPs 反而增加了 30%(因为重计算),只是显存占用少了 40-60%。一句话区分:MoD/MoE 是「该不计算的就不算了」(真省),Gradient Checkpointing 是「算完扔了,用完再算一遍」(假省、用 FLOPs 换显存)。两者完全兼容、经常一起用:训练 MoD 大模型时,同时开 MoD(省 50% FLOPs)+ Gradient Checkpointing(省 50% 显存),才能把 70B+ 的 MoD 模型塞进 8xA100 集群里训动。