Prhub

#52188 [Spec decode] Support Kimi-K3 DCP with DSpark

原始 PR 作者 wzhao18 合并时间 2026-08-18 04:08 文件变更 15 提交数 8 评论 34 代码增减 +341 / -51

执行摘要

支持 Kimi-K3 DSpark 推测解码与 DCP 并行组合

PR body 明确说明目的:"This PR adds support for running Kimi-K3 decode context parallel with DSpark with FlashinferMLA and Tokenspeed as target causal attention backend and Tokenspeed as the draft non-causal backend." 此前 vllm/config/speculative.py 直接以 ValueError 拒绝 MLA DSpark 与 decode context parallelism 组合,Kimi-K3 无法在 DCP 长上下文场景下享受推测解码加速;本 PR 移除此限制并对注意力后端、草稿输入准备、slot 映射等各层补齐 DCP 支持。

值得精读。该 PR 展示了三个可复用设计:① 热路径中把多层共享的 decode 元数据计算收敛到一次并跨层缓存(并在 review 中由维护者进一步下沉到公共基类);② DCP 下 rank-local slot 的 Triton 换算与 PAD 语义,保证草稿 KV 写入不越界;③ 以能力契约 + 启动期快速失败替代早期硬性配置拒绝,为后续后端扩展留好钩子。若要为其他 MLA 模型开启 DSpark + DCP,直接沿 _validate_dspark_dcp_supportsupports_non_causal_multi_token_dcp 两条线扩展即可。

讨论亮点

核心讨论由 reviewer GirasoleY 主导,共两条线程:

  1. 热路径效率vllm/v1/attention/backends/mla/flashinfer_mla.pyforward_mqa):GirasoleY 指出 "This is inefficient as it run for every mla forward. Let's construct the dcp related seqlen/query_start_loc/block table expansion in metadata builder, then reuse them for all MLA layers. The same metadata can be reused in combine path as well." 作者 wzhao18 回应已改为 forward pass 内只计算一次并跨层复用(即 _prepare_flattened_decode_metadata 的缓存机制)。后续 GirasoleY 又补充了一个重构 PR,将 flattened decode metadata 上移到 MLACommonDecodeMetadata 基类,最终合入版本中 FlashInferMLADecodeMetadata 直接继承该基类。

  2. FlashInferMLA 是否声明 non-causal DCP 能力(同文件第 137 行):GirasoleY 问 "Enable supports_non_causal_multi_token_dcp for flashinfer_mla as well?",wzhao18 明确回答本 PR 暂不启用,draft 模型先走 tokenspeed 后端;GirasoleY 认可并表示 "Make sense. I added a refactor PR to remove flashinfer specific changes." 这是刻意的能力边界选择,而非遗漏。

实现拆解

实现按五步拆解:

  1. 放开配置层限制vllm/config/speculative.py 删除 __post_init__ 中针对 "MLA DSpark does not currently support decode context parallelism" 的 ValueError;同步删除 vllm/models/kimi_k3/nvidia/mla.py 中相关 5 行检查,并移除 tests/transformers_utils/test_dspark_mla_config.pytest_dspark_mla_rejects_decode_context_parallelism 用例。

  2. 在 MLA backend 层建立能力契约与启动期校验vllm/model_executor/layers/attention/mla_attention.pyMLACommonMetadataBuilder 新增类变量 supports_non_causal_multi_token_dcp,并新增 _validate_dspark_dcp_support,在 __init__ 中根据 speculative_config.method == dsparkdecode_context_parallel_size > 1 时,区分 non-causal draft 与 causal multi-token 两种模式检查 backend 声明,不满足即抛 ValueError(fail fast)。vllm/v1/attention/backend.py 新增 supports_non_causal_dcp classmethod 与 validate_configurationuse_dcp 参数分支;tokenspeed_mla.py 声明 supports_non_causal_multi_token_dcp = Truevllm/v1/attention/selector.py 调整自动选择逻辑(commit 消息 "fix auto-selection to check dcp compatible"),使 DCP 场景下自动选择兼容后端。

  3. FlashInferMLA 支持 causal DCP 多 token decodevllm/v1/attention/backends/mla/flashinfer_mla.py 新增 FlashInferMLADecodeMetadata(继承 MLACommonDecodeMetadata)与 FlashInferMLAMetadataFlashInferMLAMetadataBuilder 传入 supports_dcp_with_varlen=True 并实现 _build_decodeforward_mqadcp_world_size > 1 and query_len > 1 的 causal 分支调用新方法 _prepare_flattened_decode_metadata:基于 dcp_tot_seq_lens 减去逐行偏移得到每个 query 行的全局可见前缀,再按 round-robin 换算为 rank-local seq_lens,并将展平的 block_table/seq_lens 缓存到 decode metadata 上供组内所有层复用。

  4. DFlash 草稿侧适配 DCP slot 布局vllm/v1/worker/gpu/cp_utils.py 新增 Triton kernel cp_local_slot,按 CP_INTERLEAVE 段归属 rank 计算局部 KV slot,非本 rank 位置返回 PAD_SLOT_IDvllm/v1/worker/gpu/spec_decode/dflash/speculator.py_prepare_dflash_inputs_kernel 将块号换算改为 ctx_pos // (block_size * CP_SIZE),并对 context/query slot 应用 cp_local_slot,同时 _build_draft_attn_metadatacp_size > 1 时准备 dcp_local_seq_lensvllm/v1/worker/gpu/spec_decode/dflash/cudagraph.py 在 CUDA graph 捕获路径同步传入 dcp_local_seq_lensvllm/v1/worker/gpu/spec_decode/speculator.py 的基类 _build_draft_attn_metadata 增加透传参数。

  5. 测试配套tests/v1/attention/test_mla_backends.py 新增 test_flashinfer_mla_dcp_multi_token_decode_uses_per_query_bounds(用 fake_decode 断言展开后的 seq_lens 为逐 query 前缀、block_tables 为 repeat_interleave 3 份);tests/v1/spec_decode/test_dflash_prepare_inputs.py 新增 test_prepare_dflash_inputs_excludes_rejected_context_suffix_with_dcp(覆盖 cp_rank=1/cp_size=2 下 rejected context 后缀映射为 PAD_SLOT_ID);test_flashinfer_mla_dcp.py 补充一行能力断言。

文件 模块 状态 重要度
vllm/v1/attention/backends/mla/flashinfer_mla.py MLA 后端 modified 8.6
vllm/model_executor/layers/attention/mla_attention.py 注意力抽象 modified 7.15
vllm/v1/worker/gpu/cp_utils.py 上下文并行 modified 6.55
vllm/v1/worker/gpu/spec_decode/dflash/speculator.py 推测器 modified 6.58
vllm/v1/attention/backend.py 后端基类 modified 5.83
vllm/config/speculative.py 推测配置 modified 5.28
tests/v1/attention/test_mla_backends.py MLA 测试 modified 5.98

关键符号

_prepare_flattened_decode_metadata _validate_dspark_dcp_support cp_local_slot supports_non_causal_dcp _build_draft_attn_metadata _prepare_dflash_inputs_kernel forward_mqa

关键源码片段

vllm/v1/attention/backends/mla/flashinfer_mla.py core-logic

核心实现文件。新增 FlashInferMLADecodeMetadata/FlashInferMLAMetadata 并实现 _prepare_flattened_decode_metadata,在 forward_mqa 的 causal + DCP 分支按 per-query 可见前缀生成 rank-local 展平解码元数据,并将结果缓存供整组 MLA 层复用。

def _prepare_flattened_decode_metadata(
    self,
    attn_metadata: FlashInferMLAMetadata,
    query_len: int,
    *,
    causal: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
    """准备展平后的 decode 元数据,供组内所有 MLA 层复用。    DCP causal 场景下, 每个 query 行只能看到自己位置之前的全局前缀,
    需要把多 token 的 query 块按每行可见前缀换算成 rank-local seq_lens;
    非因果 DSpark 块则简单 repeat_interleave 即可。结果缓存在 decode
    元数据上, 组内后续 layer 直接读取, 避免每层重复计算。
    """
    decode = attn_metadata.decode
    assert decode is not None
    if decode.query_len:
        # 首次调用已算过 , 直接复用本组展平结果 ( 跨层缓存 )
        assert decode.query_len == query_len
        assert decode.flattened_block_table is not None
        assert decode.flattened_seq_lens is not None
        return decode.flattened_block_table, decode.flattened_seq_lens
​
    block_table = decode.block_table.repeat_interleave(query_len, dim=0)
    if causal:
        # dcp_tot_seq_lens 是每个 request 的全局可见序列长度。
        # 对第 r 个 query 行 , 可见前缀 = 全局总长 - (query_len - 1 - r)。
        global_seq_lens = decode.dcp_tot_seq_lens
        assert global_seq_lens is not None
        offsets = torch.arange(
            query_len - 1, -1, -1,
            device=global_seq_lens.device,
            dtype=global_seq_lens.dtype,
        )
        per_query_global_lens = torch.clamp(
            (global_seq_lens.unsqueeze(1) - offsets).reshape(-1), min=0
        )
        # 把全局可见长度按 round-robin 方式分配到本 rank, 划分方式与
        # _dcp_local_seq_lens_kernel 保持一致。
        interleave = self.cp_kv_cache_interleave_size
        dcp_span = self.dcp_world_size * interleave
        remainder = torch.clamp(
            per_query_global_lens % dcp_span - self.dcp_rank * interleave,
            min=0,
            max=interleave,
        )
        seq_lens = per_query_global_lens // dcp_span * interleave + remainder
    else:
        # 非因果 DSpark 块 : 每行看到同一段上下文 , 长度一致
        seq_lens = decode.seq_lens.repeat_interleave(query_len)
​
    # 缓存到 decode 元数据 , 本组后续 layer 直接读取
    decode.flattened_block_table = block_table
    decode.flattened_seq_lens = seq_lens
    decode.query_len = query_len
    return block_table, seq_lens
vllm/v1/worker/gpu/cp_utils.py core-logic

新增 Triton kernel cp_local_slot,是 DFlash 草稿 KV 在 DCP 下正确落位的核心换算逻辑,非本 rank 位置返回 PAD_SLOT_ID。

@triton.jit
def cp_local_slot(
    positions,
    block_numbers,
    block_size,
    cp_rank,
    CP_SIZE: tl.constexpr,
    CP_INTERLEAVE: tl.constexpr,
    PAD_ID: tl.constexpr,
):
    """返回本 rank 拥有的 KV slot; 不属于本 rank 的位置返回 PAD_ID。    DCP 按 (cp_rank, cp_interleave) 粒度把每个块的槽位 round-robin 分给
    各 rank: 先按 CP_INTERLEAVE 长度切段, 段号对 CP_SIZE 取模决定归属。
    这样 DFlash 写 draft KV 时不会覆盖其他 rank 负责的槽位, 同时保证
    CUDA graph 捕获期间 slot 布局稳定。
    """
    # 位置在虚拟块 ( 跨越全部 DCP rank) 内的偏移
    block_offsets = positions % (block_size * CP_SIZE)
    if CP_SIZE == 1:
        # 无 DCP: 直接映射到物理槽位
        return block_numbers * block_size + block_offsets
    # 判断该位置是否属于当前 rank 的 interleave 段
    is_local = block_offsets // CP_INTERLEAVE % CP_SIZE == cp_rank
    # 重排为 rank-local 的连续偏移 : 先整段、再段内余量
    rounds = block_offsets // (CP_INTERLEAVE * CP_SIZE)
    remainder = block_offsets % CP_INTERLEAVE
    local_offsets = rounds * CP_INTERLEAVE + remainder
    return tl.where(is_local, block_numbers * block_size + local_offsets, PAD_ID)

评论区精华

热路径重复构造展平元数据的效率问题 性能

GirasoleY: "This is inefficient as it run for every mla forward. Let's construct the dcp related seqlen/query_start_loc/block table expansion in metadata builder, then reuse them for all MLA layers. The same metadata can be reused in combine path as well."

结论:作者改为 forward pass 内仅计算一次并缓存到 decode metadata 跨层复用;后续 GirasoleY 进一步重构,将 flattened decode metadata 上移到 MLACommonDecodeMetadata 基类。 · 已解决

FlashInferMLA 是否启用 supports_non_causal_multi_token_dcp 设计

GirasoleY: "Enable supports_non_causal_multi_token_dcp for flashinfer_mla as well?" 作者回应 "supports_non_causal_multi_token_dcp is not supported by flashinfer_mla by this PR. will let it use tokenspeed backend for the draft model for now."

结论:本 PR 刻意不为 FlashInferMLA 声明 non-causal DCP 能力,draft 后端暂由 tokenspeed 承担;GirasoleY 认可并据此做了去 flashinfer 特定化的重构。 · 已解决

风险与影响

  1. 核心路径变更forward_mqa 是 MLA decode 主路径,causal + DCP 分支新增对 dcp_tot_seq_lens 的强依赖(非 None 断言);若 batch 构建路径遗漏该字段会直接崩溃,_validate_dspark_dcp_support 只能保证 backend 声明支持,不能保证元数据一定被填充。
  2. 缓存粒度_prepare_flattened_decode_metadata 的缓存只是跨层复用,CUDA graph replay 时每个 decode step 仍会重新执行 torch.arangeclamp 等设备端小算子;相对原实现每个 layer 重复 repeat_interleave 已显著优化,但仍在每步热路径上。
  3. Triton slot 语义cp_local_slot 对非本 rank 位置返回 PAD_SLOT_ID,若 DFlash 上下文 KV 写入误覆盖会破坏他 rank 数据;单测覆盖了 rejected context suffix 与 null block 场景,但长上下文 + 多 rank 下的滑动窗口驱逐、chunked prefill 组合边界覆盖有限,e2e 仅 GSM8K。
  4. 校验语义迁移:原 speculative.py 的启动期硬拒绝改为 backend 构建期校验,新增 MLA backend 若漏声明 supports_dcp_with_varlensupports_non_causal_multi_token_dcp,会在启动时 ValueError(fail fast 属良性,但依赖后端作者遵守契约)。
  5. 能力不对称:FlashInferMLA 不支持 non-causal DCP,draft 后端只能选 tokenspeed;若用户显式指定 flashinfer 作为 draft 后端并开启 DCP 会启动失败,依赖 selector 自动选择正确降级。

用户侧影响:Kimi-K3 现在可在 DCP 长上下文场景下使用 DSpark 推测解码,默认组合 GSM8K 0.9606,FlashInferMLA 目标 + Tokenspeed 草稿 0.9613,Tokenspeed 目标 + 草稿 0.9621,精度与默认路径一致且略有提升。

系统侧影响:DFlash 的 Triton 输入准备内核与 CUDA graph 捕获路径引入 CP 维度(cp_rank/cp_size/cp_interleave),但非 DCP 场景(CP_SIZE==1)走快速分支,行为不变;注意力后端能力契约新增 supports_non_causal_multi_token_dcpsupports_non_causal_dcp,影响所有 MLA backend 的校验逻辑。

团队侧影响:确立了 MLA backend 面对 DSpark + DCP 时的声明式能力校验范式,未来扩展其他 MLA backend(如让 FlashInferMLA 支持 non-causal DCP)只需沿契约声明与 metadata 缓存两条线补齐。

核心路径变更 依赖 dcp_tot_seq_lens 断言 配置校验语义迁移 non-causal DCP 仅 tokenspeed 支持 长上下文多 rank 覆盖有限

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论