Prhub

#50327 [ModelRunnerV2] Fix scalar Mamba state update with int32 mappings

原始 PR 作者 shenoyvvarun 合并时间 2026-08-04 06:05 文件变更 2 提交数 4 评论 2 代码增减 +58 / -16

执行摘要

修复 MRV2 PP 下 Mamba 标量状态更新的 int32 索引崩溃

PR body 明确描述:ModelRunnerV2 在 pipeline-parallel 下运行混合 Mamba 模型时,非末级 pipeline rank 处理 chunked prefill 请求会在 postprocess_state 调用 index_fill_ 时抛 IndexError: index_fill_(): Expected dtype int64 for indexInputBatch.idx_mapping 有意设计为 int32 CUDA tensor,而 index_fill_ 只接受 int64 索引,因此该分支必然崩溃。作者还说明搜索过 GitHub issue 与 open PR,未发现相同报告,属于 MRV2 + PP + Mamba 场景下的首次暴露。

值得精读。本 PR 展示了如何用 Triton kernel 绕开 PyTorch op 对索引 dtype 的硬约束,并顺手处理负哨兵条目,是设备端状态更新的一种干净模式;对 MRV2、pipeline-parallel 和混合模型开发者有直接参考价值。

讨论亮点

核心讨论围绕修复方案与代码简化展开。njhill 在 approve review 中表示推送了一个简化 commit,把 num_reqs 提前计算、空 batch 提前返回、两个 kernel 调用统一到同一网格大小,作者确认接受。社区用户 subnet-dev 在 issue 中反馈在 16xH200 的 TP8+PP2 上 Kimi-K3 命中了完全相同的问题,验证了该 bug 的真实性与修复价值。claude[bot] 因 fork PR 未做自动 review,无实质内容。

实现拆解

  1. 定位根因vllm/v1/worker/gpu/model_states/mamba_hybrid.pypostprocess_state 的 else 标量分支调用 Tensor.index_fill_ 更新 num_accepted_tokens_gpu,而 index_fill_ 对索引张量的 dtype 要求是 int64,与 V1 约定的 int32 idx_mapping 不兼容。
  2. 新增 Triton kernel:定义 _fill_num_accepted_kernel,按 num_reqs 启动一维网格,每个 program 处理一行映射;读入 req_state_idx 后若为负数(-1 哨兵)直接返回,否则把常量 num_sampled 写入 num_accepted_ptr + req_state_idx,这样既处理了 PP 下被过滤的行,也避免每步在设备上分配 int64 索引转换张量。
  3. 统一控制流:在函数入口提前计算 num_reqs = idx_mapping.shape[0],空 batch 提前 return;张量与标量分支统一以 (num_reqs,) 网格调用 Triton kernel;后续 align 保存逻辑复用 num_reqs,删除原先的重复计算与嵌套 if,这是 njhill 简化 commit 的主要内容。
  4. 新增回归测试tests/v1/worker/test_mamba_hybrid_model_state.pyobject.__new__ 绕过 __init__ 构造最小状态对象,以 [2, -1, 0] 的 int32 映射同时覆盖哨兵跳过与正常写入,参数化 num_sampled=0(PP prefill 场景,归一化为 1)与 3 两个取值。
  5. 端到端验证:作者在 4 节点 32 卡 H100(PP=4, TP=8)加载 Kimi-K3 跑通多轮 chat、vision、tool calls、reasoning 与 19,569 token 长上下文,扫描全部 pipeline stage 日志确认无 index_fill_WorkerProc hit an exception 类错误。
文件 模块 状态 重要度
vllm/v1/worker/gpu/model_states/mamba_hybrid.py 模型状态 modified 7.3
tests/v1/worker/test_mamba_hybrid_model_state.py 模型状态 added 5.86

关键符号

postprocess_state _fill_num_accepted_kernel _scatter_num_accepted_kernel test_postprocess_state_scalar_with_int32_mapping

关键源码片段

vllm/v1/worker/gpu/model_states/mamba_hybrid.py data-contract

核心修复文件。`postprocess_state` 的标量分支从 `index_fill_` 改为新增的 Triton kernel `_fill_num_accepted_kernel`,解决了 int32 映射在 PP 场景下的 dtype 不匹配崩溃,并顺带简化了控制流。

# vllm/v1/worker/gpu/model_states/mamba_hybrid.py(节选)
​
    def postprocess_state(
        self,
        idx_mapping: torch.Tensor,
        num_sampled: torch.Tensor | int,
        num_computed_tokens: torch.Tensor | None = None,
    ) -> None:
        # Chunked prefill 不会采样 token,因此 num_sampled 可能为 0;
        # Mamba 把 num_accepted_tokens=1 视为非投机解码的中性值。
        num_reqs = idx_mapping.shape[0]
        if not num_reqs:
            return
​
        # idx_mapping 在 pipeline parallel 下可能含 -1 哨兵(被过滤的行),
        # kernel 直接跳过它们,避免在 host 侧做 gather。
        if not isinstance(num_sampled, int):
            _scatter_num_accepted_kernel[(num_reqs,)](
                idx_mapping, num_sampled, self.num_accepted_tokens_gpu
            )
        else:
            # 标量分支:旧实现用 index_fill_,它强制要求 int64 索引,
            # 而 InputBatch.idx_mapping 是 int32 CUDA tensor,在 PP 非末级 rank
            # 触发 IndexError。这里改用 Triton kernel 直接消费 int32 映射。
            _fill_num_accepted_kernel[(num_reqs,)](
                idx_mapping, self.num_accepted_tokens_gpu, max(num_sampled, 1)
            )
​
        # 对齐保存逻辑(spec-decode 接受后序列不再块对齐时,把运行状态
        # 保存到块对齐位置),num_computed_tokens 已持有推进后的计数。
        if (
            self._align_mode
            and num_computed_tokens is not None
            and self._mamba_ctx is not None
        ):
            self._mamba_ctx.run_fused_postprocess_align(
                num_reqs,
                self.num_accepted_tokens_gpu,
                self._mamba_state_idx_gpu,
                num_computed_tokens,
                idx_mapping,
            )
​
​
@triton.jit
def _fill_num_accepted_kernel(
    idx_mapping_ptr, # [num_reqs] batch_idx -> req_state_idx(-1 表示跳过)
    num_accepted_ptr, # [max_num_reqs]
    num_sampled,
):
    # 每个 program 处理一个请求:先读映射,跳过 -1 哨兵,再把常量
    # num_sampled 写入 num_accepted_tokens。标量值直接作为 kernel 参数传入,
    # 不需要为它单独分配 device tensor。
    row = tl.program_id(0)
    req_state_idx = tl.load(idx_mapping_ptr + row)
    if req_state_idx < 0:
        return
    tl.store(num_accepted_ptr + req_state_idx, num_sampled)
tests/v1/worker/test_mamba_hybrid_model_state.py test-coverage

新增回归测试,直接覆盖标量分支的 int32 映射与 -1 哨兵行为,参数化验证 `num_sampled=0` 归一化与普通采样值,防止该 dtype 崩溃回归。

# tests/v1/worker/test_mamba_hybrid_model_state.py(新增)@pytest.mark.skipif(not current_platform.is_cuda(), reason="Requires CUDA")
@pytest.mark.parametrize(("num_sampled", "expected_value"), [(0, 1), (3, 3)])
def test_postprocess_state_scalar_with_int32_mapping(
    num_sampled: int, expected_value: int
) -> None:
    # 用 object.__new__ 绕过 __init__,只构造被测状态对象需要的属性,
    # 让单测聚焦于 postprocess_state 的标量分支逻辑。
    state = object.__new__(MambaHybridModelState)
    state.num_accepted_tokens_gpu = torch.full(
        (4,), 9, dtype=torch.int32, device="cuda"
    )
    state._align_mode = False
    state._mamba_ctx = None
    # idx_mapping 用 int32,并包含 -1 哨兵,覆盖 PP 场景的两种典型输入:
    # 被过滤的行(跳过)和需要写入的行(写入位置 2 和 0)。
    idx_mapping = torch.tensor([2, -1, 0], dtype=torch.int32, device="cuda")
​
    state.postprocess_state(idx_mapping, num_sampled)
​
    # 期望结果:位置 0 和 2 被写入,位置 1 和 3 保持初始值 9。
    expected = torch.tensor(
        [expected_value, 9, expected_value, 9], dtype=torch.int32, device="cuda"
    )
    torch.testing.assert_close(state.num_accepted_tokens_gpu, expected)

评论区精华

index_fill_ 的 int64 索引约束与 int32 idx_mapping 的冲突 正确性

PR body 描述 traceback:`postprocess_state` 调 `index_fill_` 抛 `IndexError: index_fill_(): Expected dtype int64 for index`。V1 的 `InputBatch.idx_mapping` 有意使用 int32,而 `index_fill_` 只接受 int64,PP 非末级 rank 的 prefill 分支必然触发。

结论:新增 Triton kernel `_fill_num_accepted_kernel` 直接消费 int32 映射并跳过 -1 哨兵,避免每步分配 int64 索引转换张量。 · 已解决

njhill 的控制流简化 设计

njhill 在 approve review 中说明推送了简化 commit:提前计算 `num_reqs`、空 batch 提前返回、统一 kernel 网格大小与调用方式,并复用 `num_reqs` 到 align 分支。

结论:作者接受简化,最终合入的代码包含该提交。 · 已解决

社区用户复现确认 question

subnet-dev 在 issue 中反馈:在 16xH200 的 TP8 + PP2 上运行 Kimi-K3 命中完全相同的问题,感谢修复并期待合入。

结论:确认了问题的真实影响力与修复的必要性,无需额外改动。 · 已解决

风险与影响

回归风险:postprocess_state 控制流重排后,空 batch 提前返回与原逻辑等价;align 保存分支仅复用提前计算的 num_reqs,语义未变。平台风险:新增 Triton kernel 与回归测试均仅验证 CUDA,非 CUDA 平台(如 ROCm)没有测试保障,不过混合 Mamba 模型目前主要在 NVIDIA 上验证。性能影响:标量分支从 PyTorch op 改为 Triton launch,多一次 kernel 启动,但省去每步 int64 索引转换的设备端分配,整体为净优化。数据契约风险:idx_mapping 的 int32 与 -1 哨兵语义与既有 _scatter_num_accepted_kernel 保持一致,未来若改变 dtype 需同步更新两处 kernel。

对用户:启用 ModelRunnerV2 + pipeline parallel + Mamba 混合模型(如 Kimi-K3)的用户,prefill 后不再因状态更新崩溃,端到端服务可正常返回。对系统:状态更新热路径多一个 Triton kernel,功能等价,消除了潜在的设备端临时分配。对团队:MRV2 对 PP 场景的覆盖补齐,为后续 Mamba 混合模型在 PP 下的支持提供更稳的基础。

核心路径变更 int32 数据契约 仅 CUDA 测试覆盖 MRV2 专属路径

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论