跳转至

router_replay_utils.py — MoE 路由重放工具

文件路径: verl/utils/megatron/router_replay_utils.py

文件概述

提供 MoE 路由重放的高级工具函数,处理序列并行下的路由决策合并、分发和 PP/VPP 下的跨 rank 通信。

核心函数

1. 合并路由决策

def merge_router_topk_indices(attention_mask, input_ids, mini_layer_topk_idx_list, tf_config, vp_rank=None):
    """
    合并序列并行 rank 上记录的路由 top-k 索引。
    1. 收集所有 router 实例的记录
    2. 用 gather_from_sequence_parallel_region 合并
    3. pack/unpack 对齐到原始布局
    """

2. 设置重放数据

def set_router_replay_data(layers_topk_idx, attention_mask, tf_config, vp_rank=None):
    """
    将合并后的路由决策分发到各序列并行 rank 的 RouterReplay 实例。
    1. scatter_to_sequence_parallel_region 分发
    2. 按层索引分配到对应的 router
    """

3. PP Gather

def pp_gather(local_layers_router_map, tf_config):
    """跨 PP rank 收集路由决策,合并为全局路由图"""

4. VPP 重排

def reorder_and_merge_vpp_layers(micro_batch_tensor_list, num_microbatches, vpp_size, ...):
    """将 VPP 各 stage 的路由决策重排为连续的层顺序"""

辅助类

RouterReplayHelper

class RouterReplayHelper:
    @staticmethod
    def get_micro_batch_router_list(tf_config, vp_rank=None):
        """获取当前 micro-batch 对应的 RouterReplay 实例列表"""

    @staticmethod
    def is_r2_record_action(tf_config, vp_rank=None):
        """检查当前是否处于 RECORD 模式"""

与其他模块的关系

  • 依赖 router_replay_patch.py 的 RouterReplay 类
  • 依赖 Megatron 的序列并行和流水线并行通信原语
  • 被 MoE 模型的训练/推理 step 函数调用

小结

router_replay_utils.py 处理了 MoE 路由重放在分布式环境(TP/PP/VPP)下的复杂通信逻辑,是 MoE RLHF 训练的关键工具。