Prhub

#28714 [minimax-m3] Split 3/4: disagg K-only index-K transfer

原始 PR 作者 JustinTong0323 合并时间 2026-06-28 13:54 文件变更 7 提交数 3 评论 5 代码增减 +173 / -5

执行摘要

添加 MiniMax-M3 K-only 稀疏 KV 的 PD-disaggregation 传递

为了支持MiniMax-M3模型的PD-disaggregation,需要将稀疏KV的index-K缓冲器通过状态通道传输。Split 3/4聚焦于K-only变体(不含value),变体(含value)尚未支持并会主动抛出NotImplementedError。此PR是增量之一,确保每个PR独立可审查。

该PR设计紧凑,force_flat参数优雅地复用MLA布局,值得学习。建议在合并Split 2后尽快rebase并CI验证。对于关注disaggregation或稀疏KV的开发人员,值得精读。

讨论亮点

仅有一条来自ShangmingCai的批准评论:‘Looks good. Very clean.’,未发现争议或未解问题。

实现拆解

  1. 新增状态类型枚举:在python/sglang/srt/disaggregation/base/conn.pyStateType枚举中添加MINIMAX_INDEX_K成员,用于标识K-only稀疏索引缓冲区。

  2. 扩展setup_state_kv_args:在python/sglang/srt/disaggregation/utils.pysetup_state_kv_args函数中,通过isinstance(token_to_kv_pool, MiniMaxSparseKVPool)识别MiniMax稀疏池,若index_kv_pool存在(含value)则抛出NotImplementedError,若index_k_pool存在则调用get_index_k_state_buf_infos并追加一个MINIMAX_INDEX_K状态组件。

  3. 添加K-only传输分支:在python/sglang/srt/disaggregation/nixl/conn.pypython/sglang/srt/disaggregation/mooncake/conn.pymaybe_send_extra方法中,为StateType.MINIMAX_INDEX_K新增elif分支。该分支检查PP和TP限制(PP>1或异构TP不支持),然后调用_send_kvcache_generic并传入force_flat=True,强制使用MLA风格的平面布局(单缓冲每层),避免将K-only列表错误拆分为K/V两部分。

  4. 状态索引载荷处理:在python/sglang/srt/disaggregation/decode.pypython/sglang/srt/disaggregation/prefill.py中,在状态索引构建函数中添加StateType.MINIMAX_INDEX_K分支,复用现有的_dsa_payload()逻辑,因为索引行与主KV位于同一位置,共享页ID。

  5. 单元测试覆盖:新增test/registered/unit/disaggregation/test_minimax_sparse_disagg_state_kv_args.py,包含两个测试:test_setup_state_kv_args_single_minimax_component验证K-only池生成单个MINIMAX_INDEX_K组件;test_index_kv_pool_raises验证含value的池会正确抛出NotImplementedError

文件 模块 状态 重要度
test/registered/unit/disaggregation/test_minimax_sparse_disagg_state_kv_args.py 测试 added 7.56
python/sglang/srt/disaggregation/mooncake/conn.py 传输层 modified 6.82
python/sglang/srt/disaggregation/nixl/conn.py 传输层 modified 6.74
python/sglang/srt/disaggregation/utils.py 状态路由 modified 6.37
python/sglang/srt/disaggregation/decode.py 状态路由 modified 5.11
python/sglang/srt/disaggregation/prefill.py 状态路由 modified 5.11
python/sglang/srt/disaggregation/base/conn.py 基础类型 modified 4.58

关键符号

setup_state_kv_args maybe_send_extra _send_kvcache_generic _make_k_only_pool _make_kv_pool test_setup_state_kv_args_single_minimax_component test_index_kv_pool_raises

关键源码片段

test/registered/unit/disaggregation/test_minimax_sparse_disagg_state_kv_args.py test-coverage

新增单元测试覆盖 K-only 和 KV 两种稀疏池的状态配置,验证核心逻辑正确性。

# _make_k_only_pool: 创建仅包含 index_k 的稀疏池,用于测试
def _make_k_only_pool(start_layer: int = 0) -> MiniMaxSparseKVPool:
    # 使用 MiniMax-M3 配置:密集层 0-2,稀疏层 3-6,所有稀疏层仅 K
    dense_layer_ids = [start_layer, start_layer + 1, start_layer + 2]
    sparse_layer_ids = [start_layer + 3 + i for i in range(4)]
    end_layer = sparse_layer_ids[-1] + 1
    return MiniMaxSparseKVPool(
        size=8, page_size=4, dtype=torch.float32,
        head_num=2, head_dim=8, idx_head_dim=16,
        dense_layer_ids=dense_layer_ids, sparse_layer_ids=sparse_layer_ids,
        disable_value_sparse_layer_ids=sparse_layer_ids, # 关闭 value 层
        device="cpu", start_layer=start_layer, end_layer=end_layer,
    )# _make_kv_pool: 创建带有 index_kv(含 value)的稀疏池,期望触发 NotImplementedError
def _make_kv_pool(start_layer: int = 0) -> MiniMaxSparseKVPool:
    dense_layer_ids = [start_layer, start_layer + 1]
    sparse_layer_ids = [start_layer + 2, start_layer + 3]
    end_layer = sparse_layer_ids[-1] + 1
    return MiniMaxSparseKVPool(
        size=8, page_size=4, dtype=torch.float32,
        head_num=2, head_dim=8, idx_head_dim=16,
        dense_layer_ids=dense_layer_ids, sparse_layer_ids=sparse_layer_ids,
        disable_value_sparse_layer_ids=[], # 保留 value,index_kv_pool 不为 None
        device="cpu", start_layer=start_layer, end_layer=end_layer,
    )class TestMiniMaxSparseDisaggStateKvArgs(unittest.TestCase):
    def test_setup_state_kv_args_single_minimax_component(self):
        pool = _make_k_only_pool()
        kv_args = KVArgs()
        setup_state_kv_args(kv_args, pool)
        # 预期只有一组件,类型为 MINIMAX_INDEX_K
        self.assertEqual(kv_args.state_types, [StateType.MINIMAX_INDEX_K])
        self.assertEqual(len(kv_args.state_data_ptrs), 1)
        self.assertEqual(len(kv_args.state_data_ptrs[0]), pool.index_k_pool.layer_num)
        self.assertEqual(len(kv_args.state_item_lens[0]), pool.index_k_pool.layer_num)
​
    def test_index_kv_pool_raises(self):
        pool = _make_kv_pool()
        self.assertIsNotNone(pool.index_kv_pool)
        kv_args = KVArgs()
        # 含 value 的稀疏池应引发 NotImplementedError
        with self.assertRaises(NotImplementedError):
            setup_state_kv_args(kv_args, pool)
python/sglang/srt/disaggregation/mooncake/conn.py core-logic

在 maybe_send_extra 中添加 MINIMAX_INDEX_K 分支,使用 force_flat 参数调用 _send_kvcache_generic,实现 K-only 稀疏索引的 Mooncake 传输。

# 在 maybe_send_extra 中处理 MINIMAX_INDEX_K 状态类型
elif st == StateType.MINIMAX_INDEX_K:
    # 暂不支持 PP>1 和异构 TP(稀疏层列表压缩导致)
    if self.pp_size is not None and self.pp_size > 1:
        raise RuntimeError(
            "PD disagg: PP>1 not supported for MiniMax sparse index yet."
        )
    if (
        target_rank_registration_info is not None
        and self.attn_tp_size != target_rank_registration_info.dst_attn_tp_size
    ):
        raise RuntimeError(
            "PD disagg: heterogeneous TP not supported for MiniMax "
            "sparse index yet."
        )
    src_indices = list(indices)
    dst_indices_local = list(dst_indices)
    # 长度不一致时截断以匹配较小方
    if len(src_indices) > len(dst_indices_local):
        src_indices = src_indices[: len(dst_indices_local)]
    elif len(src_indices) < len(dst_indices_local):
        dst_indices_local = dst_indices_local[: len(src_indices)]
    rc = (
        self._send_kvcache_generic(
            mooncake_session_id=req.mooncake_session_id,
            src_data_ptrs=src_data_ptrs,
            dst_data_ptrs=dst_data_ptrs,
            item_lens=src_item_lens,
            prefill_data_indices=np.array(src_indices, dtype=np.int32),
            dst_data_indices=np.array(dst_indices_local, dtype=np.int32),
            executor=executor,
            force_flat=True, # 强制 MLA 风格平面布局,避免 KV 拆分
        ) or rc
    )
python/sglang/srt/disaggregation/utils.py dependency-wiring

在 setup_state_kv_args 中添加 MiniMaxSparseKVPool 分支,导入新池类型并根据池配置追加 MINIMAX_INDEX_K 状态组件或抛出 NotImplementedError。

# 在 setup_state_kv_args 中的新增分支
from sglang.srt.mem_cache.memory_pool import (
    DSATokenToKVPool,
    HybridLinearKVPool,
    MiniMaxSparseKVPool,
)# ... 变量初始化 ...if isinstance(token_to_kv_pool, MiniMaxSparseKVPool):
    # 如果稀疏池包含 index_kv(即含有 value 的索引),暂不支持传输
    if token_to_kv_pool.index_kv_pool is not None:
        raise NotImplementedError(
            "PD disaggregation for MiniMax sparse layers with index value "
            "(index_kv_pool) is not yet supported; only K-only sparse layers are."
        )
    # 如果只有 index_k(K-only),则获取缓冲信息并追加为 MINIMAX_INDEX_K 状态
    if token_to_kv_pool.index_k_pool is not None:
        dp, dl, il = token_to_kv_pool.get_index_k_state_buf_infos()
        append_state_component(kv_args, StateType.MINIMAX_INDEX_K, dp, dl, il)
elif hasattr(token_to_kv_pool, "get_state_buf_infos"):
    # 原有的其他池处理逻辑(如 SWA、DSA、HybridLinear 等)
    ...

评论区精华

代码审查批准 other

ShangmingCai 审查后批准,评论 'Looks good. Very clean.'

结论:PR 获得批准,无需修改。 · 已解决

风险与影响

  • 依赖风险:PR依赖#28713(Split 2)中的MiniMaxSparseKVPool,若Split 2未合并则导入失败。但CI预期会失败,计划先合并Split 2再rebase。
  • 兼容性风险:新增的MINIMAX_INDEX_K状态类型只在MiniMax-M3模型路径中激活,不影响现有模型。
  • 错误处理完善:对PP>1和异构TP使用了显式RuntimeError,避免静默数据损坏。
  • 性能影响:仅增加少量分支判断,传输路径复用现有_send_kvcache_generic框架。
  • 用户:仅MiniMax-M3用户可受益,且当前仅支持K-only变体,价值用户需等待Split 4的完整模型集成。
  • 系统:disaggregation框架新增一个状态类型,nixl和mooncake传输层各增加一个条件分支,后续扩展KV支持需修改此处。
  • 团队:代码清晰,易于审查和未来扩展。
依赖上游拆分 新状态类型 PP>1 未支持 作用域有限

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论