Grok-1通过硬编码num_selected_experts=2实现每个token仅激活2个专家,路由分三步:计算概率、top-k选专家、加权求和;8个专家物理隔离且可分片,显存增长远低于线性。
☞☞☞AI 智能聊天, 问答助手, AI 智能搜索, 多模态理解力帮你轻松跨越从0到1的创作门槛☜☜☜
要真正看懂grok-1如何用3140亿参数实现高效推理,必须打开它的源码,逐层拆解moe模块的组织逻辑、路由决策机制和专家调度方式。定位MoE核心组件:从model.py切入
进入Grok-1代码仓库根目录,直接打开model.py文件,搜索class MoELayer——这是整个混合专家架构的主干实现,位于第272行附近。这个类不是装饰性封装,而是所有专家激活、路由分发、梯度分流的实际执行入口。
紧接着向下翻,找到class Router定义(通常在MoELayer上方),它继承自hk.Module,负责为每个输入token生成专家选择概率分布。注意:Router不参与前向计算中的权重更新,只做轻量级门控决策,这是MoE保持稀疏性的第一道防线。
【num_selected_experts必须严格等于2】——查看Router初始化参数,self.num_selected_experts被硬编码为2。这意味着无论输入是什么,每个token永远只被分配给2个专家处理,这是Grok-1实现“25%参数激活率”的工程锚点,不可修改。
理解专家路由的三步计算链
在MoELayer.__call__方法内部,路由过程由三段紧密耦合的JAX操作构成:
第一步:调用self.router.compute_routing_prob(inputs, padding_mask, self.num_experts),输出routing_probs张量,形状为[batch_size, seq_len, num_experts],每个位置代表该token被分配给对应专家的概率值。
第二步:执行jax.lax.top_k(routing_probs, k=self.router.num_selected_experts),提取每行概率最高的2个索引及其数值。这一步是稀疏激活的关键开关,它直接丢弃其余6个专家的全部计算路径。
第三步:将选中的2个专家输出加权求和,权重即为top-k返回的expert_gate值;未被选中的专家输出完全不参与后续计算,内存与算力零消耗。
专家网络的物理组织方式
方法一:查看self.layer_fn参数传入方式。它是一个可调用对象,在MoELayer初始化时注入,实际指向一个封装了8个独立FFN子网络的工厂函数。每个子网络拥有完全隔离的权重矩阵,彼此不共享参数。
方法二:观察专家实例化逻辑。在__call__中,通过jnp.stack([expert(x) for expert in self.experts], axis=-2)批量调用全部专家,但随后立即用expert_index做高级索引切片——只有被top-k选中的那2列结果保留,其余6列在JAX图编译阶段就被剪枝掉。
注意:8个专家并非均匀分布在GPU上。源码中shard_activations=True配置启用后,每个专家的激活值会被自动按model_axis维度切片,跨设备分布存储,这是支撑3140亿参数模型不OOM的核心内存优化手段。
验证MoE稀疏性效果的实操检查
运行一次最小规模前向传播,在MoELayer.__call__末尾插入print(f"Activated experts: {expert_index.shape}"),你会看到输出形如Activated experts: (batch_size, seq_len, 2)——始终是2,绝不会是8或其它数字。
再检查expert_gate的sum值:对每个token位置执行jnp.sum(expert_gate, axis=-1),结果恒等于1.0。这证明路由输出是标准概率分布,且仅靠2个专家就完成了完整信息建模,没有信息坍缩。
最后对比显存占用:禁用MoE(设num_experts=1)与启用MoE(num_experts=8)两次运行nvidia-smi,你会发现后者显存增长不到前者1.3倍,而非8倍——这就是稀疏激活在真实硬件上的量化体现。


















