Prhub

#28241 [Spec v2] Use decode kernel for TRT-LLM MHA draft extend

原始 PR 作者 hanming-lu 合并时间 2026-06-16 07:57 文件变更 1 提交数 2 评论 3 代码增减 +4 / -1

执行摘要

TRT-LLM MHA draft extend 使用 decode kernel

让 TRT-LLM MHA draft extend 阶段使用 decode kernel 以获得更好的性能。之前曾尝试过但遇到问题(issue #24863),本次重新落地。

建议合并,但需确保 CI 中覆盖了 TRT-LLM MHA 后端 + speculative decoding 的测试用例。同时,未来可考虑增加对应的回归测试。

讨论亮点

无review评论。仅在PR body和合并者评论中提到这是重新落地,之前尝试失败过(issue #24863)。

实现拆解

  1. 修改 python/sglang/srt/layers/attention/trtllm_mha_backend.pyforward_extend 方法的分支条件:将原来的只检查 is_target_verify() 扩展为同时检查 is_target_verify()is_draft_extend_v2()
  2. 当条件满足时,调用 flashinfer.decode.trtllm_batch_decode_with_kv_cache(decode kernel)而非 flashinfer.prefill.trtllm_batch_context_with_kv_cache(prefill kernel)。
  3. 移除了之前用于灰度测试的环境变量 SGLANG_TRTLLM_MHA_DRAFT_EXTEND_V2_DECODE(在 python/sglang/srt/environ.py 中定义)。
文件 模块 状态 重要度
python/sglang/srt/layers/attention/trtllm_mha_backend.py 注意力层 modified 6.11

关键符号

forward_extend

关键源码片段

python/sglang/srt/layers/attention/trtllm_mha_backend.py core-logic

核心变更文件,修改了 forward_extend 方法的分支条件,使 draft_extend_v2 走 decode kernel 路径。

# 在 forward_extend 方法中,原条件只检查 is_target_verify(),
# 现在增加 is_draft_extend_v2(),使 draft_extend_v2 阶段也使用 decode kernel。
if (
    forward_batch.forward_mode.is_target_verify()
    or forward_batch.forward_mode.is_draft_extend_v2()
):
    # 使用 decode kernel 执行 attention,替代 prefill kernel
    o = flashinfer.decode.trtllm_batch_decode_with_kv_cache(
        query=q,
        kv_cache=kv_cache,
        workspace_buffer=self.workspace_buffer,
        block_tables=page_table,
        seq_lens=self.forward_metadata.cache_seqlens_int32,
        max_seq_len=self.max_context_len,
        bmm1_scale=bmm1_scale,
        bmm2_scale=bmm2_scale,
        window_left=layer.sliding_window_size,
        sinks=attention_sink,
        skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(),
        out_dtype=self.q_data_type,
        q_len_per_req=self.forward_metadata.max_seq_len_q,
    )
else:
    # 原有 prefill kernel 路径保持不变
    o = flashinfer.prefill.trtllm_batch_context_with_kv_cache(...)

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

该变更直接将 draft_extend_v2 的 kernel 路径从 prefill 切换到 decode,可能存在以下风险:

  1. 如果 decode kernel 在某些场景下(如特殊序列长度、head 配置)不支持或表现异常,可能导致推理错误或性能回退。
  2. 缺少回归测试用例,未覆盖该分支的验证。
  3. 由于移除了环境变量,无法通过配置快速回退到旧行为,需要代码回滚。
    但鉴于该路径在之前版本中已被灰度测试(通过环境变量),风险相对可控。

影响范围:仅影响使用 TRT-LLM MHA 后端的 speculative decoding 场景,特别是 draft_extend_v2 阶段。用户无感知,但推理性能可能改善。由于移除了环境变量,运维灵活性略有降低。

核心路径变更 缺少测试覆盖

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论