执行摘要
- 一句话:添加MiniMax-M3 K-only稀疏KV的PD-disaggregation传递
- 推荐动作:该PR设计紧凑,
force_flat参数优雅地复用MLA布局,值得学习。建议在合并Split 2后尽快rebase并CI验证。对于关注disaggregation或稀疏KV的开发人员,值得精读。
功能与动机
为了支持MiniMax-M3模型的PD-disaggregation,需要将稀疏KV的index-K缓冲器通过状态通道传输。Split 3/4聚焦于K-only变体(不含value),变体(含value)尚未支持并会主动抛出NotImplementedError。此PR是增量之一,确保每个PR独立可审查。
实现拆解
-
新增状态类型枚举:在python/sglang/srt/disaggregation/base/conn.py的StateType枚举中添加MINIMAX_INDEX_K成员,用于标识K-only稀疏索引缓冲区。
-
扩展setup_state_kv_args:在python/sglang/srt/disaggregation/utils.py的setup_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状态组件。
-
添加K-only传输分支:在python/sglang/srt/disaggregation/nixl/conn.py和python/sglang/srt/disaggregation/mooncake/conn.py的maybe_send_extra方法中,为StateType.MINIMAX_INDEX_K新增elif分支。该分支检查PP和TP限制(PP>1或异构TP不支持),然后调用_send_kvcache_generic并传入force_flat=True,强制使用MLA风格的平面布局(单缓冲每层),避免将K-only列表错误拆分为K/V两部分。
-
状态索引载荷处理:在python/sglang/srt/disaggregation/decode.py和python/sglang/srt/disaggregation/prefill.py中,在状态索引构建函数中添加StateType.MINIMAX_INDEX_K分支,复用现有的_dsa_payload()逻辑,因为索引行与主KV位于同一位置,共享页ID。
-
单元测试覆盖:新增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(模块 测试;类别 test;类型 test-coverage;符号 _make_k_only_pool, _make_kv_pool, TestMiniMaxSparseDisaggStateKvArgs, test_setup_state_kv_args_single_minimax_component): 新增单元测试覆盖K-only和KV两种稀疏池的状态配置,验证核心逻辑正确性。
python/sglang/srt/disaggregation/mooncake/conn.py(模块 传输层;类别 source;类型 core-logic): 在maybe_send_extra中添加MINIMAX_INDEX_K分支,使用force_flat参数调用_send_kvcache_generic,实现K-only稀疏索引的Mooncake传输。
python/sglang/srt/disaggregation/nixl/conn.py(模块 传输层;类别 source;类型 core-logic): 在maybe_send_extra中添加MINIMAX_INDEX_K分支,使用force_flat参数调用_send_kvcache_generic,实现K-only稀疏索引的NIXL传输。
python/sglang/srt/disaggregation/utils.py(模块 状态路由;类别 source;类型 dependency-wiring): 在setup_state_kv_args中添加MiniMaxSparseKVPool分支,导入新池类型并根据池配置追加MINIMAX_INDEX_K状态组件或抛出NotImplementedError。
python/sglang/srt/disaggregation/decode.py(模块 状态路由;类别 source;类型 core-logic): 在状态索引构建中添加MINIMAX_INDEX_K分支,复用DSA载荷逻辑,因为索引行与主KV位于同一位置。
python/sglang/srt/disaggregation/prefill.py(模块 状态路由;类别 source;类型 core-logic): 在状态索引构建中添加MINIMAX_INDEX_K分支,复用DSA载荷逻辑,与decode保持一致。
python/sglang/srt/disaggregation/base/conn.py(模块 基础类型;类别 source;类型 core-logic): 新增StateType.MINIMAX_INDEX_K枚举成员,作为K-only稀疏索引状态的标识。
关键符号: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
新增单元测试覆盖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
在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
在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 等)
...
评论区精华
仅有一条来自ShangmingCai的批准评论:‘Looks good. Very clean.’,未发现争议或未解问题。
- 代码审查批准 (other): 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未支持, 作用域有限
关联脉络
- PR #28713 [minimax-m3] Split 2/4: mem-cache / HiCache / sparse KV pool: 此PR依赖Split 2提供MiniMaxSparseKVPool。
- PR #27944 MiniMax-M3 full feature (original large PR): 此PR是原MiniMax-M3大PR的拆分部分。
参与讨论