Prhub

#45019 [NIXL][Mamba] Add Mamba1 support to NIXL P/D disaggregation

原始 PR 作者 Josephasafg 合并时间 2026-06-25 20:50 文件变更 5 提交数 18 评论 13 代码增减 +104 / -35

执行摘要

NIXL 连接器支持 Mamba1 模型 P/D 分离

PR body明确目标是“Add Mamba1 hybrid-model support (e.g. Jamba) to the NIXL connector's conv-state transfer for prefill/decode disaggregation”。此前NIXL连接器仅支持Mamba2/GDN的conv-state传输,Mamba1模型(如Jamba)无法利用P/D分离能力。

值得精读。该PR展示了如何通过泛化数据结构而非添加分支条件来扩展已有功能,设计简洁。对于维护kv-connector的工程师,理解 MambaConvSplitInfo 的偏移计算逻辑对后续修改至关重要。

讨论亮点
  • 模型选择:NickLucche 建议使用更小的Mamba1模型进行CI测试,从 ai21labs/AI21-Jamba2-Mini(52B)改为 ai21labs/AI21-Jamba2-3B
  • 内联建议:ZhanqiuHu 建议将 _ssm_regions_per_layer 属性内联到 _compute_desc_ids 中,作者接受并实施。
  • 回归验证:ZhanqiuHu 要求检查Mamba2/GDN是否仍工作,作者用 Qwen3.5-0.8B 验证1P4D和4P1D,结果在预期范围内。

实现拆解

  1. 泛化数据结构:在 ssm_conv_transfer_utils.py 中将 MambaConvSplitInfo.local_proj_dimstuple[int, int, int] 改为 tuple[int, ...]proj_bytes 对应变为 tuple[int, ...]
  2. 重构偏移计算local_conv_offsetsremote_conv_offsets 从硬编码三段解包改为基于 proj_bytes 的循环累加,支持任意子投影数。
  3. 调整描述符ID计算:在 base_worker.py_compute_desc_ids 中,将 num_ssm_regions 从固定值 len(self.block_len_per_layer) * 4 改为根据 len(self._conv_decomp.local_conv_offsets) + 1 动态计算。
  4. 添加Mamba1测试:在单元测试 test_nixl_connector_hma.py 中增加 jamba_mini_tp1/4/8 参数化用例;在集成测试配置 config_sweep_accuracy_test.sh 中添加 ai21labs/AI21-Jamba2-3B 运行项;在 test_accuracy.py 中添加精度基线。
  5. 回归验证:用 Qwen3.5-0.8B(GDN Mamba2)在1P4D和4P1D下验证,精度与预期一致(0.33±0.05),确保Mamba2/GDN不受影响。
文件 模块 状态 重要度
vllm/distributed/kv_transfer/kv_connector/v1/ssm_conv_transfer_utils.py SSM 传输工具 modified 7.72
vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py NIXL 工作器 modified 6.46
tests/v1/kv_connector/unit/test_nixl_connector_hma.py 单元测试 modified 5.74
tests/v1/kv_connector/nixl_integration/config_sweep_accuracy_test.sh 集成测试 modified 3.82
tests/v1/kv_connector/nixl_integration/test_accuracy.py 精度测试 modified 3.11

关键符号

MambaConvSplitInfo.__init__ MambaConvSplitInfo.proj_bytes MambaConvSplitInfo.local_conv_offsets MambaConvSplitInfo.remote_conv_offsets NixlBaseConnectorWorker._compute_desc_ids test_derive_mamba_conv_split

关键源码片段

vllm/distributed/kv_transfer/kv_connector/v1/ssm_conv_transfer_utils.py core-logic

核心变更文件,泛化 Mamba1 卷积状态子投影数据结构

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Mamba conv-state 子投影分解工具,用于 NIXL 传输。
   支持 Mamba1(单子投影)、Mamba2(x/B/C)和 GDN(Q/K/V)。
"""import math
from dataclasses import dataclass
import torch
from vllm.model_executor.layers.mamba.mamba_utils import is_conv_state_dim_first
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
from vllm.v1.kv_cache_interface import MambaSpec
​
​
@dataclass(frozen=True)
class MambaConvSplitInfo:
    """每个 rank 的卷积状态子投影字节大小。
       使用 DS 连续布局(dim, state_len),子投影在内存中连续。       内存布局在一个 page 内:
         Mamba1: |---- x ----|  (单个子投影,无分解)
         Mamba2: |-- x --|- B -|- C -|  (B == C)
         GDN:    |- Q -|- K -|-- V --|  (dim(Q)==dim(K), V 可能不同)
    """
    conv_rows: int # conv_kernel - 1(通常为 3)
    # 每个 rank 各子投影的列数,
    # Mamba1 为 1 个值,Mamba2/GDN 为 3 个值
    local_proj_dims: tuple[int, ...]
    conv_dtype_size: int # 每元素字节数(如 float16 为 2)
    ssm_sizes: tuple[int, int] # (conv_state_bytes, ssm_state_bytes)
​
    @property
    def local_conv_dim(self) -> int:
        """当前 rank 的总卷积列数"""
        return sum(self.local_proj_dims)
​
    @property
    def proj_bytes(self) -> tuple[int, ...]:
        """返回各子投影的字节大小(本地 rank)"""
        row_bytes = self.conv_rows * self.conv_dtype_size
        return tuple(d * row_bytes for d in self.local_proj_dims)
​
    @property
    def local_conv_offsets(self) -> list[tuple[int, int]]:
        """返回 (byte_offset, byte_size) 列表,供本地描述符注册"""
        offsets: list[tuple[int, int]] = []
        offset = 0
        for size in self.proj_bytes:
            offsets.append((offset, size))
            offset += size
        return offsets
​
    def remote_conv_offsets(
        self, local_rank_offset: int, tp_ratio: int
    ) -> list[tuple[int, int]]:
        """返回本 D rank 在 P page 内的子投影切片 (byte_offset, byte_size)。
           tp_ratio 标识 P/D 的 tensor 并行大小关系。
        """
        offsets: list[tuple[int, int]] = []
        if tp_ratio >= 1:
            # D_TP >= P_TP;P page 更大,D 读取其切片
            remote_base = 0
            for size in self.proj_bytes:
                offsets.append((remote_base + local_rank_offset * size, size))
                remote_base += size * tp_ratio
        else:
            # P_TP > D_TP;P page 更小,D 需要按比例缩放
            abs_ratio = -tp_ratio
            remote_base = 0
            for size in self.proj_bytes:
                remote_size = size // abs_ratio
                offsets.append((remote_base, remote_size))
                remote_base += remote_size
        return offsets

评论区精华

选择更小的 Mamba1 测试模型 设计

NickLucche 建议使用更小的 Mamba1 模型进行 CI 测试,避免使用 52B 的 Jamba2-Mini。作者最终采用 Jamba2-3B。

结论:使用 ai21labs/AI21-Jamba2-3B 作为 CI 测试模型。 · 已解决

内联 _ssm_regions_per_layer 设计

ZhanqiuHu 建议将 _ssm_regions_per_layer 属性内联到 _compute_desc_ids 中,以减少状态。

结论:作者采纳,最终版本中直接计算 ssm_regions_per_layer 局部变量。 · 已解决

回归验证 Mamba2/GDN 正确性

ZhanqiuHu 要求确认 Mamba2/GDN 仍然正常工作。作者用 Qwen3.5-0.8B 在 1P4D 和 4P1D 配置上验证,精度在预期范围内。

结论:回归通过。 · 已解决

风险与影响

核心数据结构从固定三投影改为可变长度,若 local_proj_dims 长度不符合预期(如Mamba2/GDN意外传入1个维度),可能导致偏移计算错误。但测试覆盖了三种类型且回归验证通过。性能影响可忽略,循环只多走1-3次。兼容性要求设置 VLLM_SSM_CONV_STATE_LAYOUT=DS 环境变量,未设置时新功能不会启用。

用户侧:支持Mamba1混合模型(如Jamba)的P/D分离,降低显存占用。系统侧:改动集中在两个核心源文件,影响范围有限。团队侧:后续添加更多SSM变体(如Mamba3)时,泛化结构可直接复用。

核心路径变更 环境变量依赖

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论