Mixture-of-Depths: Dynamically allocating compute in transformer-based language models
David Raposo, Sam Ritter, Blake Richards, Timothy Lillicrap, Peter Conway Humphreys, Adam Santoro
cs.LG, cs.CL
2024-04-03
DeepMind 用 expert-choice 的 top-k 路由,让隔层 Transformer 块只处理 12.5% 的 token。isoFLOP 下损失不差于稠密模型,训练步进可快约 60%,前向 FLOPs 最多少一半,解码几乎不掉点。
标准 Transformer 对序列里每个位置、每一层都付同样的 FLOPs。不是每个 token 都需要同等深度:功能词、已确定的续写,和要做长程推理的位置,难度差一截。条件计算想把算力花在需要的地方,但动态计算图和未知张量尺寸跟现有硬件不对付。MoE 用路由在专家之间切分计算,总 FLOPs 并不下降。
Mixture-of-Depths 要的是一张静态图、事先定死的容量,同时让 token 在深度上走不同条路:有的层做完整的自注意力和 MLP,有的层直接走残差,省掉那一层的计算。
每个 MoD 块的自注意力和 MLP 只接收 k 个 token,k 是容量,训练前定死。路由器给每个 token 打一个标量分,expert-choice 取出 top-k 送进块,其余原样残差传下去。选 expert-choice 是因为负载天然均衡,不必再加平衡损失,而且分高的 token 一定进得去,重要位置不容易被丢掉。路由器权重乘在块输出上,让路由决策走在梯度路径上。
容量小于序列长度,自注意力的 query-key 乘法从 T² 降到 k²。k=T/2 时这一项 FLOPs 是原来的 25%。总计算量在训练前就可预测,硬件友好;哪些 token 被选中则随上下文变。
Expert-choice 的 top-k 是非因果的:一个 token 能不能进 top-k,取决于后面 token 的分数,自回归解码时看不到未来。两条补法。一是给路由器加二元交叉熵,用 top-k 选择当标签,把分数推向 0.5 两侧,语言模型损失大约多 0.2% 到 0.3%。二是另训一个小 MLP 预测该 token 会不会进 top-k,主损失不受影响。这个辅助任务很简单,准确率很快到 99%。
和早期退出不同:MoD 允许 token 跳过中间层,之后再和走过那些层的 token 做注意力。和 MoE 也不同:选项是「算」或「不算」,不是在多个 MLP 专家里挑一个,而且路由同时作用于注意力和 MLP。
序列长度 2048,batch 128。小预算 6e18 FLOPs 上扫超参:隔一层做路由、容量 12.5%(2048 里只留 256 个 token 进块)最好。每层都路由会伤性能;随机路由远差于学习路由和稠密基线。把容量再压过 12.5%,开始掉点。
isoFLOP 曲线被 MoD 往下、往右推:最优点损失更低,参数量更大。一个 2.2 亿参数的 MoD,略好于同样 2.2 亿的 isoFLOP 最优稠密基线,训练步进快约 60%,墙钟训练时间接近。更大预算 2e19 和 1e20、模型从 6 千万到 30 亿,趋势还在:存在既更快又更低损失的 MoD。经验规则是,把 MoD 的每步 FLOPs 调到和 isoFLOP 最优稠密模型一样,再优先加深度而不是宽度。
自回归评测在 25.6 万条、约 5 亿 token 的留出集上。训练用非因果 top-k、解码改预测器,性能几乎不掉。MoD 还可以和 MoE 拼成 MoDE:先做 MoD 再进 MoE(staged),或把 no-op 当成一个专家(integrated)。两者都能把 MoD 的增益叠在 MoE 上;integrated 比单纯把专家容量改小、靠丢 token 当残差更好,因为模型会明确学会选「不算」。
这是 2024 年就把「按 token 动态深度」做成静态图的一次干净演示。对要降推理延迟的人,isoFLOP 最优点旁边那些更小的 MoD,用更少每步 FLOPs 打平稠密最优损失,解码步进可以快一半以上。对要压训练墙钟的人,这篇的设定是训练 FLOPs 和墙钟对齐稠密模型,换来更大、更好的网络,或者同样质量下更快的 step。
它和 MoE 正交,可以叠。路由同时改写谁更新、谁能被看见,等于连 KV 集合也在做稀疏选择,这和只疏 MLP 的 MoE 不是同一件事。
论文报告的是训练损失和步进速度,没有 MMLU 这类下游表。把它当成已经验证的通用能力提升,过了。
非因果 top-k 是硬伤,解码必须靠辅助损失或预测器近似。99% 的预测准确率听起来高,长序列上 1% 的路由错误会不会在深度上累积,论文没有拆开看。最优配置「隔层、12.5% 容量」来自 6e18 的小预算扫描,更大模型是否仍成立,只在 isoFLOP 曲线上间接看到,没有单独的容量扫描。
没有公开的标准下游评测,也没有和 CoLT5、early-exit、token merging 在同一解码延迟轴上比。内存节省、KV cache 变小是作者预期,正文说没有系统测。训练用固定超参家族,序列 2048,外推到长上下文和专家并行集群,还是空白。