跳转至

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 训练的关键基础设施。