router_replay_patch.py — MoE 路由重放补丁¶
文件路径: verl/utils/megatron/router_replay_patch.py
文件概述¶
实现 MoE (Mixture of Experts) 模型的路由重放 (Router Replay) 机制。在 RLHF 训练中,rollout(推理)阶段记录的路由决策可以在训练阶段重放,确保专家分配一致性。
背景知识¶
MoE 模型中,Router 决定每个 token 分配给哪些 Expert。在 RLHF 中: 1. Rollout 阶段生成 response 时,Router 做出路由决策 2. 训练阶段计算 loss 时,需要用相同的路由决策(否则梯度不正确) 3. Router Replay 机制记录并重放这些决策
核心类¶
RouterReplay — 路由重放管理¶
class RouterReplay:
router_instances = [] # 所有 MoE 层的 router 实例
@staticmethod
def set_replay_data(all_layers_topk_indices):
"""设置所有层的重放数据"""
for i, router in enumerate(RouterReplay.router_instances):
router.set_target_indices(all_layers_topk_indices[i])
@staticmethod
def get_recorded_data():
"""获取所有层记录的路由决策"""
return [router.get_recorded_indices() for router in RouterReplay.router_instances]
RouterReplayAction — 动作枚举¶
class RouterReplayAction(Enum):
RECORD = "record" # 记录模式
REPLAY_FORWARD = "replay_forward" # 前向重放
REPLAY_BACKWARD = "replay_backward" # 反向重放
补丁应用¶
def apply_router_replay_patch():
"""通过 monkey-patch 修改 Megatron TopKRouter 的行为"""
# 1. 给 TransformerConfig 添加 enable_routing_replay 属性
# 2. 修改 TopKRouter.__init__ 为每个 router 创建 RouterReplay 实例
# 3. 修改 TopKRouter.routing 方法支持 RECORD / REPLAY
# 4. 修改 MoEAlltoAllTokenDispatcher.preprocess 处理重复索引
TopKRouter.routing = patched_routing
工作流程¶
Rollout 阶段 (RECORD):
Token → Router → topk_indices (记录) → Expert 计算 → Output
训练阶段 (REPLAY_FORWARD):
Token → Router → 使用记录的 topk_indices → Expert 计算 → Loss
训练阶段 (REPLAY_BACKWARD):
反向传播也使用记录的 topk_indices
与其他模块的关系¶
- 被
router_replay_utils.py用来管理全局路由状态 - 被 Megatron MoE worker 在训练时调用
- 依赖
megatron.core.transformer.moe.router.TopKRouter
小结¶
router_replay_patch.py 通过 monkey-patch 技术扩展了 Megatron 的 MoE Router,使其支持路由决策的记录和重放,是 MoE 模型 RLHF 训练的关键基础设施。