AI Pulse

用RTX 3090从零训8天,446M MoE模型逼近GPT-2 medium

用RTX 3090从零训8天,446M MoE模型逼近GPT-2 medium

一个446M参数的MoE模型,用一张RTX 3090从零训练,耗时不到8天。最终测试损失3.253928,比作者之前训练的所有模型都好。它也超过了OpenAI GPT-2 small(124M参数)。和GPT-2 medium(345M参数)比,还差一点,但已经很接近。

模型有6个专家,每个token激活2个。作者的说法是“GPT-2 small加上MoE改造”。指令微调测试里它排第3,赢了作者所有其他模型。输给的是OpenAI的两个模型,连124M的GPT-2 small也没赢过。

MoE的核心是稀疏激活:每次推理只动用部分参数。内存占用和同等大小的密集模型一样,但快得多。用作者的话说,推理速度接近小模型,知识储备接近大模型。

这项技术已经被用起来了。Claude和ChatGPT被广泛传闻是MoE,DeepSeek、Kimi K3等开源大模型明确在用。

作者特意强调,“专家”不按领域分工。没有谁管语法、谁管事实。六个专家学到的是一些难以理解的语言模式,它们服务于语言建模这个整体目标。

负载不均衡:MoE的典型坑

MoE有个突出难点:负载均衡。只按训练损失优化,路由器会把大部分token塞给少数专家。其他专家被“饿死”,整体性能跟着掉。

作者第一次没加负载均衡机制,跑了2天、3.2B tokens,测试损失3.501952。图表显示,某些专家被强烈偏好,另一些几乎没被用到。

解法是引入辅助损失,公式来自Switch Transformers:loss = α·N·Σ(f_i·P_i)。f_i是分配给第i个专家的token比例,P_i是路由器给第i个专家的平均权重。这个损失会把路由器的选择拉向均匀。

辅助损失的强度也要调。作者跑了一小时的训练来搜索α,最终选0.005。完美平衡时,Switch Transformers的辅助损失值为1(单活跃专家)。作者的模型同时激活2个专家,所以理想值是2。

实现细节与对照

路由器是一个线性层,把上下文向量映射到专家数量。top-k操作选出每个token对应的专家。scatter_函数再把分数写进一个全-∞的tensor,生成掩码后的logits。

如果只有一个活跃专家,softmax权重恒为1,路由器梯度为零,没法训练。作者的代码里对“少于两个活跃专家”的情况直接抛异常。

作者把自己的实现跟Hugging Face上的Mixtral源码做了对比。MoE处理方式大方向一致,唯一差别在辅助损失计算。Mixtral把不同层的相同专家索引当成同一个专家来算,作者觉得这“有点奇怪”。

作者还让ChatGPT、Claude、Kimi K3三个LLM审了代码。它们的结论一致:“一个相当标准的GPT-2之上的MoE实现,负载均衡改编自Switch Transformers。”

作者参考了四篇论文:1991年的《Adaptive Mixtures of Local Experts》、2017年的《Outrageously Large Neural Networks》、2020年的《GShard》、2021年的《Switch Transformers》。

下一步

GPT-2的前馈网络(FFN)参数是注意力机制的两倍。MoE改造的主要目标就是它:把FFN复制多份,用路由器分发。

除了token选择路由,还有“专家选择”方案。它反转逻辑:不为每个token挑top-k专家,改为给每个专家挑top-k token。这样从数学上保证负载均衡。作者的实现用的是token选择方案,靠辅助损失解决均衡问题。

训练里还有一个他没完全想清楚的地方:Switch Transformers的f_i项不可微。作者目前的处理是把它当常数。

作者的计划很直接:用完全相同的计算量,训练一个446M参数的密集模型。他想知道MoE相比密集模型到底领先多少。

阅读原文
📚 相关主题 大语言模型开源

订阅 AI Pulse

每天 08:00 · 12:30 · 18:30 · 23:50 更新