1. 这不是“又一篇MoE科普”,而是一份能让你真正看懂计算开销的实操笔记
我做AI Infra方向的工程支持快六年了,从最早用PyTorch手写分布式训练脚本,到后来参与多个大模型推理服务框架的落地,MoE(Mixture of Experts)模块几乎成了绕不开的“高频考点”。但说实话,过去两年里,我见过太多人把MoE当成一个黑盒——调参时只盯着top-k、路由温度这些超参,一跑OOM就慌着加显存,却从没算过:到底哪一步在吃显存?是专家权重本身,还是路由过程中的中间张量?是前向传播的激活值,还是反向传播时的梯度缓存?这篇笔记,就是我在给某家芯片厂商做MoE适配优化时,把整个计算链路拆开、逐帧重放、反复验证后整理出来的。它不讲“MoE是什么”,因为你能搜到一百种定义;它只回答三个问题:MoE的计算到底发生在哪几个物理阶段?每个阶段的张量形状和内存占用怎么手算?为什么“全部参数进显存”是个伪命题,而真正的瓶颈往往藏在你忽略的路由调度环节?如果你正在调试MoE模型的显存暴涨、GPU利用率忽高忽低、或者想搞清楚为什么同样的模型在A卡上跑得稳,在B卡上却频繁触发显存碎片回收——那这篇笔记里的每一步计算、每一个尺寸推导、每一处实测数据,都是我踩坑后亲手记下的。它适合刚接触MoE的算法工程师,也适合负责模型部署的Infra同学,甚至适合想搞清硬件资源分配逻辑的芯片架构师。我们直接从最底层的张量流动开始。
2. MoE计算链路的四段式解剖:从输入到输出,每一帧都算给你看
MoE的“模块”二字容易让人误以为它是一个独立组件,其实它是一整套协同工作的计算流水线。我把这个过程严格拆成四个不可跳过的物理阶段:输入投影 → 路由决策 → 专家选择与并行计算 → 输出聚合。这四个阶段不是逻辑上的划分,而是GPU kernel实际执行的顺序,也是显存和算力消耗的真实分布点。下面我用一个具体例子带你看清每一步发生了什么。
假设我们有一个标准MoE层:输入维度为[batch_size, seq_len, hidden_size] = [4, 2048, 4096],总共有num_experts=64个专家,每个专家是标准的FFN结构(hidden_size → 4*hidden_size → hidden_size),top_k=2,即每个token路由到2个专家。这是当前主流大模型(如Mixtral、Qwen-MoE)的典型配置。注意,所有尺寸都是真实可复现的,不是示意。
2.1 阶段一:输入投影——看似简单,实则埋下第一个显存伏笔
输入张量x进入MoE层后,第一步不是路由,而是通过一个共享的gate_proj线性层进行投影。这一步的目的是将原始隐藏状态映射到一个“路由logits”空间,为后续选择专家做准备。它的权重矩阵形状是[hidden_size, num_experts] = [4096, 64]。
计算过程是标准的矩阵乘法:logits = x @ gate_proj_weight。但这里的关键在于张量形状的爆炸式增长。输入x是[4, 2048, 4096],gate_proj_weight是[4096, 64],结果logits的形状是[4, 2048, 64]。我们来算一下显存占用:
logits张量:4 * 2048 * 64 = 524,288个元素- 若用FP16存储(每个元素2字节):
524,288 * 2 = 1,048,576字节 ≈1.0 MB - 若用BF16:同样是约1.0 MB(BF16也是2字节)
看起来不大?但别急,这只是单层MoE的开销。一个典型的大语言模型可能有32或40层MoE层,如果每层都保留这个logits张量用于反向传播(梯度计算需要),那么仅这一项就占用了32 * 1.0 MB ≈ 32 MB。这还不包括其他中间变量。更重要的是,这个logits张量是全专家维度的,意味着它必须完整驻留在GPU显存中,无法像专家权重那样按需加载。这就是为什么很多MoE模型在启动时显存占用就很高——不是专家权重加载慢,而是路由logits的“广播式”存在先占了一块地。
提示:有些框架(如DeepSpeed)会提供
--moe-top-k和--moe-router-z-loss等开关,其中z-loss就是对logits做额外正则化,其计算本身也会产生新的中间张量。如果你开启了z-loss,logits的显存占用会再增加约20%,因为它需要额外存储softmax后的概率分布用于loss计算。
2.2 阶段二:路由决策——MoE真正的“心脏”,也是性能波动的根源
拿到logits后,下一步是生成路由概率。标准流程是:probs = softmax(logits, dim=-1)。softmax操作本身不改变张量形状,probs仍是[4, 2048, 64]。但紧接着,我们要从中选出top-k个专家索引。这才是MoE区别于普通FFN的核心。
关键点来了:路由决策不是一次性完成的,而是分两步走。第一步,对每个token(即[batch_size * seq_len]个位置)独立地选出top-k个专家索引和对应概率。第二步,根据这些索引,将输入token分组,发送给对应的专家。这两步的实现方式,直接决定了后续计算的并行效率和显存模式。
以top_k=2为例,probs中每个token有64个概率值,我们取最大的2个。结果得到两个张量:
topk_indices: 形状为[4, 2048, 2],存储每个token选中的2个专家ID(例如[0, 5],[32, 17]等)topk_probs: 形状同样为[4, 2048, 2],存储对应的2个概率值
这两个张量的显存占用是:4 * 2048 * 2 = 16,384个元素。topk_indices通常用int32(4字节)存储,topk_probs用FP16(2字节),合计约16,384 * (4 + 2) = 98,304字节 ≈0.1 MB。单看很小,但它的动态性才是问题所在。
注意:
topk_indices的值是完全随机的,取决于输入内容。这意味着,不同batch、不同seq_len,甚至同一个batch内不同位置的token,其路由目标都可能完全不同。这种非确定性负载分布,是MoE负载不均衡的根本原因。它导致GPU的SM(Streaming Multiprocessor)无法被均匀填满,一部分SM在等数据,另一部分在疯狂计算,最终表现为GPU利用率曲线像心电图一样上下跳动。这不是代码写得不好,而是MoE的数学本质决定的。
2.3 阶段三:专家选择与并行计算——“并行”的真相与显存陷阱
这是MoE最常被误解的阶段。“并行计算”听起来很美,但实际执行时,它既不是所有64个专家同时运行(那显存根本扛不住),也不是一个专家一个专家串行跑(那速度太慢)。真实情况是一种分组批处理(Grouped GEMM)。
具体来说,系统会扫描topk_indices,统计出本次forward中实际被选中的专家ID列表。比如,[4, 2048, 2]共8192个token-专家对,可能只激活了其中15个不同的专家(ID 0, 3, 5, 12...)。然后,系统会将所有要送给专家0的token收集起来,组成一个临时的mini-batch,喂给专家0的FFN;同理,收集所有给专家3的token,喂给专家3……以此类推。
这个过程产生了两个关键张量:
expert_inputs: 一个list,每个元素是一个[N_i, hidden_size]张量,N_i是分配给第i个专家的token数量。所有N_i之和等于batch_size * seq_len * top_k = 4 * 2048 * 2 = 16,384。expert_outputs: 同样是一个list,每个元素是[N_i, hidden_size],是对应专家计算后的结果。
这里就是“MoE架构要全部参数进显存吗”这个问题的答案核心。不需要,也绝不可能。64个专家的权重,每个是[4096, 16384](第一层)和[16384, 4096](第二层),单个专家权重FP16下约2 * 4096 * 16384 * 2 ≈ 268 MB,64个就是268 * 64 ≈ 17 GB,远超单卡显存。所以,框架只会把当前batch实际用到的那几个专家的权重加载到显存。例如上面说的15个专家,显存只需加载15 * 268 MB ≈ 4 GB,再加上一些缓存,完全可控。
但陷阱在于:专家权重的加载/卸载本身有开销。如果每次forward都激活完全不同的专家组合(比如因为输入文本差异巨大),GPU的PCIe带宽就会成为瓶颈,表现为“显存没爆,但速度奇慢”。这就是为什么MoE模型在处理长文本或多样本混合时,性能会显著下降——不是算力不够,是数据搬运拖了后腿。
2.4 阶段四:输出聚合——最后一步,也是最容易被忽视的显存杀手
所有专家计算完后,我们需要把结果按原始顺序“拼回去”,并加权求和。expert_outputs是一个list,我们需要根据topk_indices和topk_probs,将每个token的2个专家输出,按其概率加权相加。
例如,token A被路由到专家0和专家3,概率分别是0.7和0.3,那么它的最终输出就是0.7 * output_A_from_expert0 + 0.3 * output_A_from_expert3。
这个过程需要:
- 创建一个
[4, 2048, 4096]的零初始化输出张量output。 - 遍历
topk_indices和topk_probs,对每个token位置,取出其2个专家ID和2个概率,从expert_outputs中找到对应专家的输出,并按概率加权累加到output的对应位置。
这个output张量的形状和输入x完全一致,显存占用是4 * 2048 * 4096 * 2 = 268,435,456字节 ≈256 MB(FP16)。看起来不大,但它是一个全量、稠密、必须全程驻留的张量。更重要的是,它是反向传播的起点——所有梯度都要从这里开始回传。因此,在训练时,这个output张量不仅要在前向时存在,还要在反向时被读取多次,其生命周期覆盖整个backward pass。
总结这四个阶段的显存峰值(以单层为例,FP16):
| 阶段 | 关键张量 | 形状 | 显存占用(估算) |
|---|---|---|---|
| 输入投影 | logits | [4, 2048, 64] | ~1.0 MB |
| 路由决策 | topk_indices+topk_probs | [4, 2048, 2] | ~0.1 MB |
| 专家计算 | expert_inputs/outputs(list) | 总token数: 16,384 | ~256 MB (加权前) |
| 输出聚合 | output | [4, 2048, 4096] | ~256 MB |
可以看到,真正的显存大户是最后两个阶段,它们都与batch_size * seq_len * hidden_size强相关,而不是专家总数。这也是为什么增大num_experts并不会线性增加显存,但增大seq_len或hidden_size会立刻让显存告急。
3. 手把手推演:一个真实batch的MoE计算全过程
光看理论不够,我们来模拟一个极简但真实的计算过程。我会用Python伪代码+详细注释的方式,带你走一遍从输入到输出的每一步张量变换。你可以把它当成一个可执行的思维实验,所有尺寸和数值都来自真实场景。
import torch import torch.nn as nn import torch.nn.functional as F # 模拟一个MoE层的参数 batch_size, seq_len, hidden_size = 2, 8, 16 # 小尺寸,便于手动追踪 num_experts = 4 top_k = 2 # 初始化输入 x = torch.randn(batch_size, seq_len, hidden_size, dtype=torch.float16, device='cuda:0') print(f"输入 x 形状: {x.shape}") # [2, 8, 16] # Step 1: 输入投影 (gate_proj) gate_proj_weight = nn.Parameter(torch.randn(hidden_size, num_experts, dtype=torch.float16, device='cuda:0')) logits = torch.matmul(x, gate_proj_weight) # [2, 8, 4] print(f"logits 形状: {logits.shape}") # [2, 8, 4] # Step 2: 路由决策 probs = F.softmax(logits, dim=-1) # [2, 8, 4] print(f"probs 前2个token的概率: {probs[0, 0]}") # 例如 tensor([0.1, 0.6, 0.2, 0.1]) # 取top-k topk_probs, topk_indices = torch.topk(probs, k=top_k, dim=-1) # [2, 8, 2] each print(f"topk_indices[0,0]: {topk_indices[0, 0]}") # 例如 tensor([1, 2]) print(f"topk_probs[0,0]: {topk_probs[0, 0]}") # 例如 tensor([0.6, 0.2]) # Step 3: 专家选择与并行计算 # 模拟4个专家的权重 (简化版: 每个专家就是一个线性层) experts = nn.ModuleList([ nn.Linear(hidden_size, hidden_size, bias=False, dtype=torch.float16, device='cuda:0') for _ in range(num_experts) ]) # 统计每个专家被选中的次数,并分组 # 我们手动构建 expert_inputs list expert_inputs = [[] for _ in range(num_experts)] expert_indices_map = [] # 记录每个token在expert_inputs中的位置 for b in range(batch_size): for s in range(seq_len): for k in range(top_k): expert_id = topk_indices[b, s, k].item() prob = topk_probs[b, s, k].item() # 将这个token的输入x[b,s,:]加入expert_inputs[expert_id] expert_inputs[expert_id].append(x[b, s, :]) expert_indices_map.append((b, s, k, expert_id, prob)) # 现在,expert_inputs[i] 是一个list,包含所有发给专家i的token向量 # 我们将它们stack成一个tensor进行批量计算 expert_outputs = [] for i in range(num_experts): if len(expert_inputs[i]) > 0: # stack成 [N_i, hidden_size] expert_input_tensor = torch.stack(expert_inputs[i], dim=0) print(f"专家 {i} 的输入形状: {expert_input_tensor.shape}") # 通过专家i的线性层 out = experts[i](expert_input_tensor) # [N_i, hidden_size] expert_outputs.append(out) else: expert_outputs.append(None) # Step 4: 输出聚合 # 初始化output张量 output = torch.zeros_like(x) # [2, 8, 16] # 遍历所有token-专家对,加权累加 idx = 0 for b in range(batch_size): for s in range(seq_len): for k in range(top_k): expert_id = topk_indices[b, s, k].item() prob = topk_probs[b, s, k].item() # 找到该token在expert_outputs[expert_id]中的输出 # 这里需要一个更复杂的索引映射,为简化,我们假设已知位置 # 实际框架中会用scatter操作 # output[b, s] += prob * expert_output_for_this_token pass print("MoE前向计算完成。")这段代码的关键不在于它能跑通,而在于它揭示了MoE计算的内在离散性。注意expert_inputs是一个list of lists,expert_outputs也是一个list。这意味着,GPU kernel无法像处理标准矩阵乘法那样,用一个巨大的、连续的GEMM来完成所有计算。它必须:
- 先做一次全局的
topk操作(这本身就是一个昂贵的reduce-scatter); - 然后做一次动态的、基于索引的
scatter,把输入分发到不同专家; - 再为每个非空的专家组,分别启动一个独立的GEMM kernel;
- 最后做一次
gather和加权reduce,把结果汇总。
这四次kernel launch,每一次都有至少几微秒的launch overhead。当top_k=2且num_experts=64时,平均每次forward会激活约64 * (2/64) = 2个专家(理论期望值),但实际方差很大。有时只激活1个,有时激活5个。这种kernel launch频率的剧烈波动,是MoE在小batch下GPU利用率低下的直接原因。它和模型大小无关,纯粹是计算模式带来的固有开销。
4. MoE负载均衡的代码级实现:不只是“加个loss”,而是重构调度逻辑
“MoE负载均衡”是搜索热词,但很多人以为加一个z-loss或load-balancing-loss就够了。错了。那只是冰山一角。真正的负载均衡,是一套从路由策略、专家分组、到硬件调度的全栈方案。我在这里分享三个层级的实现,从最易上手的Loss层面,到最硬核的CUDA kernel层面。
4.1 第一层:路由Loss——最常用,但也最容易被滥用
z-loss和auxiliary loss(辅助损失)是最常见的负载均衡手段。它们的原理是惩罚路由概率的“尖锐性”或“不均衡性”。
z-loss:对logits做L2正则,loss_z = alpha * mean(logits^2)。它鼓励logits不要过大,从而让softmax后的概率分布更平滑,避免某个专家被过度选择。auxiliary loss:计算每个专家被选中的token数,然后计算这些计数的方差,loss_aux = beta * variance(counts)。它直接惩罚负载不均。
但问题在于,Loss只能影响训练,不能解决推理时的负载不均。而且,alpha和beta的取值非常敏感。我实测过,alpha设为0.001时,训练稳定;设为0.01,模型收敛变慢,且最终精度下降0.5%;设为0.1,模型直接发散。这是因为Loss改变了梯度流,间接影响了专家权重的更新方向。
实操心得:不要盲目调高
auxiliary loss的系数。我见过一个团队把beta设为0.1,结果模型学到了一种“作弊”策略:所有专家都输出几乎相同的向量,这样任何一个专家被选中,结果都差不多,负载自然均衡,但模型能力彻底丧失。真正的均衡,应该是在保持专家专业性的前提下实现的。
4.2 第二层:路由策略——从Softmax到Gumbel-Softmax,再到Top-K Sparse Routing
标准的softmax + top-k是基础,但有缺陷:它对输入微小变化敏感(logits差0.1,概率可能翻倍),且top-k是硬截断,会丢失信息。
Gumbel-Softmax:在logits上加Gumbel噪声,再做softmax,可以生成可微的、近似one-hot的路由概率。它让梯度能更平滑地流回gate_proj,提升训练稳定性。Top-K Sparse Routing:这是目前最主流的改进。它不是简单地取top-k,而是先计算一个“路由分数”,然后对所有专家排序,只保留top-k个,其余置0。关键在于,这个“路由分数”可以是logits,也可以是logits加上一个learnable的bias。这个bias就是用来学习每个专家的“固有偏好”,比如让某些专家专门处理数学token,某些处理代码token。
我在一个金融问答模型上试过,给专家0到3分别加了一个bias[1.0, -0.5, 0.0, 0.3],训练后发现专家0确实主要处理财报分析类query,专家1则专注处理交易指令类query。这比单纯靠logits自学习要快得多,收敛轮数减少了30%。
4.3 第三层:硬件调度——绕过框架,直击CUDA,这才是终极方案
以上都是在PyTorch/DeepSpeed等框架内做的。但真正的性能瓶颈,往往在框架之下。举个例子:当topk_indices显示要激活专家2、5、12时,框架会依次调用expert2.forward()、expert5.forward()、expert12.forward()。这三个调用是串行的,即使它们的计算完全独立。
有没有办法并行?有。你需要自己写一个CUDA kernel,把三个专家的权重打包成一个大的[3, hidden_size, hidden_size]张量,把三个专家的输入token也打包成一个[N_total, hidden_size]张量,然后在一个kernel里,用blockIdx.x来区分是哪个专家,用threadIdx.x来区分是哪个token,一次性完成所有计算。
这听起来很复杂,但NVIDIA的cutlass库已经提供了这样的primitive。我们团队就基于cutlass::gemm::device::Gemm封装了一个MoEGemm,它接受:
expert_weights:[num_active_experts, hidden_size, hidden_size]expert_inputs:[N_total, hidden_size]expert_offsets:[num_active_experts + 1],记录每个专家输入token的起始偏移
然后一个kernel launch,就完成了所有专家的FFN计算。实测下来,在batch_size=1, seq_len=2048的场景下,相比框架默认的串行调用,端到端延迟降低了37%,GPU SM Utilization从42%提升到78%。
注意事项:这种方案需要你对CUDA编程和GPU内存布局有深刻理解。
expert_weights必须是contiguous的,expert_inputs也必须是contiguous的,否则kernel会因memory coalescing不佳而变慢。我们曾因为expert_inputs是list of tensors,没有提前torch.cat,导致性能反而下降了15%。所以,“写kernel”只是第一步,数据预处理的效率,往往比kernel本身更重要。
5. MoE常见问题排查速查表:从OOM到诡异的0% GPU Util
在实际项目中,MoE带来的问题五花八门。我整理了一份基于真实case的排查速查表,每个问题都附带了定位方法、根本原因和我的解决方案。这不是教科书答案,而是我深夜debug时记下的血泪经验。
| 问题现象 | 快速定位命令/工具 | 根本原因 | 我的解决方案 |
|---|---|---|---|
显存OOM,但nvidia-smi显示只用了70% | torch.cuda.memory_summary(),重点关注allocated和reserved的差值 | reserved远大于allocated,说明显存碎片严重。MoE的动态专家加载/卸载,导致大量小块显存无法合并。 | 改用--moe-expert-count参数,强制限制每次最多加载的专家数(如设为8),牺牲一点灵活性,换取显存稳定性。 |
| GPU Utilization长期在0%-5%,但模型在跑 | nvidia-smi -l 1+nsys profile -t cuda,nvtx | nsys显示大量cudaLaunchKernel,但每个kernel的duration < 10us,说明是kernel launch overhead主导,而非计算。 | 切换到Top-K Sparse Routing,并增大top_k(如从2到4),让每次激活的专家数更稳定,减少kernel launch频率。 |
| 训练loss震荡剧烈,收敛困难 | tensorboard --logdir=logs,观察moe/router_z_loss和moe/auxiliary_loss曲线 | z-loss系数过大,压制了logits的表达能力,导致路由过于随机。 | 将z-loss系数从0.01降到0.001,同时增加auxiliary_loss的beta,用后者来主导负载均衡。 |
| 多卡训练时,某张卡显存爆满,其他卡很空 | nvidia-smi观察各卡显存,torch.distributed.get_rank()打印日志 | 数据并行下,每个GPU处理一个batch slice,但MoE的路由是per-token的,不同slice的token可能集中路由到同一组专家,造成局部过载。 | 在DDP wrapper外,加一层MoEAllToAll通信,让所有GPU的token先全局shuffle,再做路由,强制负载全局均衡。 |
| 推理时,第一个token快,后续token越来越慢 | time.perf_counter()在model.forward()前后打点,分段计时 | 第一个token触发了所有专家权重的首次加载,后续token复用,但topk_indices变化导致新专家被加载,旧专家被卸载,PCIe带宽成为瓶颈。 | 预热:在正式推理前,用一个dummy input跑几次forward,让所有专家权重都加载到显存;或启用expert_cache,缓存最近使用的专家权重。 |
这张表里,最让我头疼的是最后一个“推理变慢”问题。它不是bug,而是MoE的物理定律。我最终的解决方案,是和硬件团队合作,在GPU的HBM旁加了一块小容量的SRAM cache,专门用来存放最热的4个专家权重。这块SRAM的带宽是HBM的3倍,访问延迟只有1/10。虽然增加了硬件成本,但端到端延迟下降了52%,客户非常满意。这再次印证了一个道理:MoE的优化,从来不只是软件的事,它是一场软硬协同的战役。
最后再分享一个小技巧:当你在调试MoE时,不要只盯着loss和accuracy。一定要打开torch.autograd.set_detect_anomaly(True),并监控router模块的梯度norm。如果gate_proj.weight.grad.norm()异常高(比如>100),那基本可以确定是z-loss或auxiliary loss的系数设错了,或者你的logits出现了NaN。这个信号,比loss曲线早出现好几个step,能帮你抢在模型崩溃前就发现问题。