Prhub

#30378 [DSA] Re-enable fused top-k v2 for MTP: clamp padded-row seq_lens to >= 0

原始 PR 作者 DarkSharpness 合并时间 2026-07-08 04:44 文件变更 6 提交数 2 评论 8 代码增减 +57 / -28

执行摘要

修复 DP MTP 中 fused top-k v2 的非法内存访问

PR body 明确说明:"Removes its TEMP allow_topk_v2 gate and fixes the root cause of the GLM 5.2 MTP illegal memory access"。根因分析指出:在 DP attention 下,draft-extend-v2 CUDA graph replay 中 idle-companion / DP-padded rows 携带的 seq_len fill value (1) 小于 qo_len,导致 expanded seq_lens 为负(如 -4),被 top-k v2 kernel 作为 uint32 读取后变为 ~4e9 token 长度,产生非法地址。此问题导致 DP MTP 在首个请求后立即 crash。

值得精读。作者对 root cause 的定位方法(CUDA-GDB register 分析、instrumented replay)是高质量调试范例;clamp 策略明确、文档契约化的设计值得借鉴;渐进式的 PR 演化(先 gate 再修复)展示了安全的迭代模式。

讨论亮点

PR body 中作者对根因进行了详尽的 CUDA-GDB 调试分析,但 review 评论无实质性讨论。审核者 Fridge003 直接批准。

实现拆解

  1. Clamp 负值 seq_lens:在 python/sglang/srt/layers/attention/triton_ops/dsa_metadata.py_fused_dsa_draft_extend_metadata_kernelpython/sglang/srt/layers/attention/triton_ops/pad.pyseqlens_expand_kernel 中,对 expanded seq_lens 添加 tl.maximum(..., 0) 操作,确保任何负值被截断为 0。0 长度会被 kernel 视为无 token,走 trivial all-(-1) 安全路径。
  2. 移除 allow_topk_v2 gate:在 python/sglang/srt/layers/attention/dsa_backend.pyDSAIndexerMetadata 中删除 allow_topk_v2 字段,并在 get_indexer_metadata 中移除对 is_target_verify / is_draft_extend_v2 的屏蔽;在 python/sglang/srt/layers/attention/dsa/dsa_topk_backend.pytopk_transform 中移除 allow_topk_v2 参数和条件判断,使 fused top-k v2 对符合 shape 条件的 PAGED 场景(包括 spec verify / draft-extend)统一生效。
  3. 强化文档契约:在 python/sglang/jit_kernel/dsv4/topk.pyplan_topk_v2topk_transform_512_v2 函数中添加详细 docstring,明确要求 seq_lens 必须非负,并解释负值的后果。在 _topk_transform_v2_paged 的 docstring 中也补充了非负契约。
  4. 修复单元测试回归:在 test/registered/kernels/test_dsa_indexer.py_run_fused_topk_backend_equivalence_test 中,构建 mock DSAMetadata 时调用 plan_topk_v2 预计算 topk_v2_plan,避免 fused v2 dispatch 的 assertion 失败。
文件 模块 状态 重要度
python/sglang/srt/layers/attention/dsa/dsa_topk_backend.py TopK 调度 modified 6.39
python/sglang/srt/layers/attention/dsa_backend.py 索引器元数据 modified 5.81
python/sglang/srt/layers/attention/triton_ops/dsa_metadata.py Triton 内核 modified 3.49
python/sglang/srt/layers/attention/triton_ops/pad.py Triton 内核 modified 3.49
python/sglang/jit_kernel/dsv4/topk.py JIT 内核 modified 5.17
test/registered/kernels/test_dsa_indexer.py 单元测试 modified 4.13

关键符号

topk_transform _topk_transform_v2_paged _fused_dsa_draft_extend_metadata_kernel seqlens_expand_kernel plan_topk_v2 topk_transform_512_v2 get_indexer_metadata _run_fused_topk_backend_equivalence_test

关键源码片段

python/sglang/srt/layers/attention/dsa/dsa_topk_backend.py core-logic

核心调度:移除了 allow_topk_v2 gate,更新注释为 spec verify / draft-extend 也走 v2 路径,并强化了 lengths 非负契约

# dsa_topk_backend.py 中 topk_transform 方法的 dispatch 逻辑
# 现在 fused top-k v2 对所有符合条件的 PAGED 场景(包括 spec verify / draft-extend)开放
def topk_transform(
    self,
    logits: torch.Tensor,
    lengths: torch.Tensor,
    topk: int,
    topk_transform_method: TopkTransformMethod,
    attn_metadata,
    ...
) -> torch.Tensor:
    if not envs.SGLANG_DSA_FUSE_TOPK.get() or force_unfused_topk:
        return self.topk_func(logits, lengths, topk, row_starts=row_starts)
​
    # 关键分支:fused top-k v2 路径(移除 allow_topk_v2 gate)
    if (
        envs.SGLANG_OPT_USE_TOPK_V2.get()
        and topk_transform_method == TopkTransformMethod.PAGED
        and row_starts is None
        and batch_idx_list is None
        and 0 < topk <= 2048
        and lengths.shape[0] == logits.shape[0] == attn_metadata.real_page_table.shape[0]
    ):
        return _topk_transform_v2_paged(logits, lengths, topk, attn_metadata)
​
    # 回退到 legacy 路径(需要 page_table_1)
    assert attn_metadata.page_table_1 is not None
    # ... ( 后续代码省略 )
python/sglang/srt/layers/attention/triton_ops/dsa_metadata.py infrastructure

在 fused_dsa_draft_extend_metadata kernel 中添加 clamp,防止负 seq_lens

# dsa_metadata.py 中 _fused_dsa_draft_extend_metadata_kernel 的 clamp
# 防止负 seq_lens 被下游 uint32 kernel 识别为超大值
expanded_seq = base_seq - qo_len_for_row + local_off + 1
expanded_seq = tl.maximum(expanded_seq, 0) # Clamp to >= 0: DP-padded rows 的安全兜底
expanded_seq = tl.where(mask_e, expanded_seq, 0)
dsa_seq = tl.minimum(expanded_seq, dsa_index_topk)
dsa_cu = tl.cumsum(dsa_seq, 0)

评论区精华

无实质性讨论 other

PR 主要由作者在 body 中详细分析,review 评论无额外讨论。

结论:直接批准,无需修改。 · 已解决

风险与影响

风险较低,因修复经过本地(4x B200 DP 配置)和 CI rerun 验证(GLM5.2 MTP 测试通过)。主要风险点:clamp 到 0 可能掩盖其他错误的 seq_lens 负值来源(但逻辑上合理的 padding 应该得到 0);另外重新启用 fused top-k v2 可能在其他未测试的模型或配置中触发类似问题,但相同 shape 条件的 DP MTP 已覆盖。单元测试的修复确保合并后主分支测试通过。

影响范围集中于使用 DP attention 且启用 MTP(multi-token prediction)的场景,主要为 GLM 5.2 和 DeepSeek-V3.2 等模型。受益用户无需再遭遇 illegal memory access 崩溃,且 spec decode 性能恢复至 fused top-k v2 水平(avg_spec_accept_length ~4.0)。对非 DP 或非 MTP 场景无影响。团队维护工作量降低(移除临时 gate)。

CUDA graph 填充值依赖 DP padding 假设

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论