Prhub

#44214 [RL Infra][FlashInfer] Enable router replay output from FlashInfer monolithic MoE kernel

原始 PR 作者 xuanyu-mistral 合并时间 2026-07-21 07:45 文件变更 11 提交数 26 评论 41 代码增减 +1099 / -12

执行摘要

FlashInfer monolithic MoE 内核添加 routing replay 捕获

强化学习训练需要捕获每层专家的路由决策以进行专家负载均衡(EPLB)。原有的RoutedExpertsCapturer仅支持modular kernel路径,而FlashInfer的融合monolithic MoE内核将路由隐藏在kernel内部,导致路由决策丢失。此PR通过向kernel传递routing_replay_out张量,使其写出所选专家ID,从而统一两种路径的捕获能力。

值得精读。该PR展示了如何在融合kernel中嵌入回调的设计模式(预分配缓冲区 + 后分发回调),以及通过基类接口定义与子类职责分离的实践。GPU Model Runner中的快速失败保护比静默返回全零更安全。测试代码为特定硬件支持的kernel编写端到端测试提供了样板,值得测试团队参考。

讨论亮点

Review中主要讨论点:

  • 支持范围:aoshen02询问是否仅DeepSeek V3模型支持,作者回复routing_replay_out工作在所有路由方法上,随后删除了之前对DeepSeek V3的限制。
  • 缓冲区分配优化:aoshen02建议减少运行时分配开销,作者将缓冲区改为预分配(在set_capture_fn时分配,仅当尺寸不足时重新分配),避免热路径每次分配。
  • 捕获函数设置重复:aoshen02指出monolithic路径下可能同时设置routerfused_experts的capture_fn,作者通过条件分支区分,确保只设置一次。
  • 性能验证:aoshen02要求对比四种配置的性能和logprob差异,作者提供了详细数据,证明显著无回归。
  • 代码风格:aoshen02要求保留被误删的注释,作者恢复。

实现拆解

实现分为以下步骤:

  1. 基类契约定义:在modular_kernel.pyFusedMoEExpertsMonolithic中新增supports_routing_replay_captureset_capture_fn_maybe_make_routing_replay_buffer_maybe_dispatch_routing_replay方法。set_capture_fn预分配一个(max_num_tokens, experts_per_token)的int16缓冲区,避免热路径重复分配。_maybe_make_routing_replay_buffer在capture_fn非空时返回缓冲区(否则返回None),_maybe_dispatch_routing_replay在kernel执行后调用回调。
  2. 子类启用:在trtllm_fp8_moe.pytrtllm_bf16_moe.pytrtllm_nvfp4_moe.pytrtllm_mxint4_moe.pytrtllm_mxfp4_moe.py中覆盖supports_routing_replay_capture返回True,并在各自的apply方法中调用_maybe_make_routing_replay_buffer获取缓冲区,将缓冲区作为routing_replay_out传入FlashInfer kernel,kernel填充后调用_maybe_dispatch_routing_replay触发回调。
  3. GPU Model Runner绑定:修改gpu_model_runner.py_bind_routed_experts_capturer方法,增加monolithic路径检测:如果quant_method.is_monolithic为True,则检查fused_experts是否为FusedMoEExpertsMonolithic实例且supports_routing_replay_capture()返回True,否则抛出ValueError;若通过则调用fused_experts.set_capture_fn。对于modular路径保持原有行为。
  4. 辅助修改:在routed_experts_capturer.py中添加topk维度断言;在flashinfer_mxint4_moe.py中传递routing_replay_out参数。
  5. 测试配套:新增tests/kernels/moe/test_routed_experts_capture_monolithic.py,包含三组Blackwell-only kernel测试,覆盖BF16、FP8、NVFP4 monolithic kernel,使用TopK、GroupedTopK、Renormalize三种路由方法,并验证禁用捕获时跳过分配。同时扩展现有tests/model_executor/test_routed_experts_capture.py,增加monolithic不支持时的错误测试和dummy fixtures适配。
文件 模块 状态 重要度
vllm/model_executor/layers/fused_moe/modular_kernel.py MoE 核心 modified 8.49
tests/kernels/moe/test_routed_experts_capture_monolithic.py MoE 测试 added 8.14
tests/model_executor/test_routed_experts_capture.py MoE 测试 modified 7.37
vllm/v1/worker/gpu_model_runner.py GPU 运行器 modified 7.33
vllm/model_executor/layers/fused_moe/experts/trtllm_fp8_moe.py MoE 专家 modified 6.94

关键符号

supports_routing_replay_capture set_capture_fn _maybe_make_routing_replay_buffer _maybe_dispatch_routing_replay _bind_routed_experts_capturer _capture_fn

关键源码片段

vllm/model_executor/layers/fused_moe/modular_kernel.py data-contract

核心基类变更,新增 capture API 定义了 monolithic kernel 支持路由捕获的契约。

# vllm/model_executor/layers/fused_moe/modular_kernel.py (FusedMoEExpertsMonolithic)# 新增类属性
routing_replay_capture_fn: Callable[[torch.Tensor], None] | None = None
_routing_replay_buffer: torch.Tensor | None = Nonedef supports_routing_replay_capture(self) -> bool:
    """子类如果支持路由回放捕获(例如 FlashInfer 提供了``routing_replay_out``),应覆盖此方法返回 True。"""
    return Falsedef set_capture_fn(
    self,
    capture_fn: Callable[[torch.Tensor], None] | None,
) -> None:
    """设置回调函数,并在非 None 时预分配缓冲区。"""
    self.routing_replay_capture_fn = capture_fn
    if capture_fn is None:
        self._routing_replay_buffer = None
        return
    # 预先分配 (max_num_tokens, experts_per_token) 的 int16 缓冲区,
    # 避免在每次 apply 时重新分配。
    self._routing_replay_buffer = torch.empty(
        (self.moe_config.max_num_tokens, self.moe_config.experts_per_token),
        dtype=torch.int16,
        device=self.moe_config.device,
    )def _maybe_make_routing_replay_buffer(
    self,
    num_tokens: int,
    device: torch.device,
) -> torch.Tensor | None:
    """如果绑定了回调,返回预分配的缓冲区(大小不足时抛出 ValueError);否则返回 None。"""
    if self.routing_replay_capture_fn is None:
        return None
    buf = self._routing_replay_buffer
    assert buf is not None
    if buf.shape[0] < num_tokens or buf.device != device:
        raise ValueError(
            "Routing replay buffer 已初始化为 "
            f"{buf.shape[0]} 个 token 位于 {buf.device},但 kernel 收到了 "
            f"{num_tokens} 个 token 位于 {device}。"
        )
    return bufdef _maybe_dispatch_routing_replay(
    self,
    routing_replay_out: torch.Tensor | None,
    num_tokens: int,
) -> None:
    """kernel 执行后,如果存在回调且缓冲区非空,调用回调传递前 num_tokens 行。"""
    if routing_replay_out is None or self.routing_replay_capture_fn is None:
        return
    self.routing_replay_capture_fn(routing_replay_out[:num_tokens])
vllm/v1/worker/gpu_model_runner.py data-contract

修改 _bind_routed_experts_capturer,增加 monolithic 路径检测与 fail-fast 保护。

# vllm/v1/worker/gpu_model_runner.py (GPUModelRunner._bind_routed_experts_capturer)def _bind_routed_experts_capturer(self, capturer: RoutedExpertsCapturer) -> None:
    from vllm.model_executor.layers.fused_moe.layer import MoERunner
    from vllm.model_executor.layers.fused_moe.modular_kernel import (
        FusedMoEExpertsMonolithic,
    )
    from vllm.model_executor.layers.fused_moe.router.base_router import (
        BaseRouter,
    )
​
    for module in self.model.modules():
        if not isinstance(module, MoERunner):
            continue
        layer_id = module.layer_id
​
        def _capture_fn(topk_ids, _layer_id=layer_id, _capturer=capturer):
            _capturer.capture(_layer_id, topk_ids)
​
        quant_method = module._quant_method
        moe_kernel = getattr(quant_method, "moe_kernel", None)
        impl = getattr(moe_kernel, "impl", None)
        fused_experts = getattr(impl, "fused_experts", None)
​
        if quant_method.is_monolithic:
            # 对于 monolithic kernel,需要 fused_experts 支持 routing replay
            if not (
                isinstance(fused_experts, FusedMoEExpertsMonolithic)
                and fused_experts.supports_routing_replay_capture()
            ):
                raise ValueError(
                    "--enable-return-routed-experts is not supported with "
                    f"monolithic MoE kernel {type(fused_experts).__name__}; "
                    "routed expert IDs would be silently all-zero."
                )
            fused_experts.set_capture_fn(_capture_fn)
        elif isinstance(module.router, BaseRouter):
            # 对于 modular kernel,通过 router 设置回调
            module.router.set_capture_fn(_capture_fn)

评论区精华

Monolithic kernel 支持的 routing method 范围 question

aoshen02 询问 `trtllm_bf16_moe.py` 中新增 `supports_routing_replay_capture` 是否仅支持 DeepSeek V3 的路由方法,并质疑其他 kernel 的支持范围。

结论:作者确认 `routing_replay_out` 工作在所有路由方法上,随后删除了之前对 DeepSeek V3 的限制。 · 已解决

运行时减少缓冲区分配开销 性能

aoshen02 建议减少每次 apply 时分配 `routing_replay_out` 缓冲区的开销,改为静态 / 持久化分配。

结论:作者在 `modular_kernel.py` 中将缓冲区改为在 `set_capture_fn` 时预分配,仅当尺寸不足时才重新分配,避免热路径分配。 · 已解决

Monolithic 路径下 capture_fn 设置避免重复 正确性

aoshen02 指出在 `_bind_routed_experts_capturer` 中,对 monolithic kernel 可能同时通过 `module.router.set_capture_fn` 和 `fused_experts.set_capture_fn` 设置回调,导致重复或冲突。

结论:作者通过条件分支区分:monolithic 时只设置 `fused_experts.set_capture_fn`,modular 时仍通过 `router.set_capture_fn`,确保只设置一次。 · 已解决

移除代码注释的合理性 style

aoshen02 要求不要移除 `modular_kernel.py` 中 `apply` 方法前的注释,认为其提供了重要的参数说明。

结论:作者恢复被误删的注释。 · 已解决

风险与影响

  • 仅Blackwell GPU:TRTLLM fused MoE kernel需要SM100+,测试会自动跳过不兼容平台,生产环境使用其他GPU无影响。
  • fail-fast保护:若用户在monolithic kernel不支持capture时启用--enable-return-routed-experts_bind_routed_experts_capturer会抛出ValueError,避免返回全零占用位。
  • 运行时开销:额外分配一个int16 (max_num_tokens, topk)缓冲区,持久化后开销可接受;回调仅在capture_fn非空时执行,默认情况下不影响推理性能。
  • 回归风险:修改涉及多个文件但数据流清晰,测试覆盖了主要路径和错误路径,回归概率较低。
  • 用户:RL实训用户可直接从monolithic kernel获得路由数据,无需强制使用modular kernel。非RL用户不受影响(需显式开启--enable-return-routed-experts)。
  • 系统:不影响默认推理路径;新增代码仅在启用capture时执行,推理性能无退化。
  • 团队:提供标准化的集成模板,后续新增monolithic kernel只需覆盖一个方法并插入两行调用,降低开发成本。
Blackwell-only GPU 需显式开启 feature 运行时断言开销

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论