# PR #28241 完整报告

- 仓库：`sgl-project/sglang`
- 标题：[Spec v2] Use decode kernel for TRT-LLM MHA draft extend
- 合并时间：2026-06-16 07:57
- 原文链接：http://prhub.com.cn/sgl-project/sglang/pull/28241

---

# 执行摘要

- 一句话：TRT-LLM MHA draft extend 使用 decode kernel
- 推荐动作：建议合并，但需确保 CI 中覆盖了 TRT-LLM MHA 后端 + speculative decoding 的测试用例。同时，未来可考虑增加对应的回归测试。

# 功能与动机

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

# 实现拆解

1. 修改 `python/sglang/srt/layers/attention/trtllm_mha_backend.py` 中 `forward_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`（模块 注意力层；类别 source；类型 core-logic；符号 forward_extend）: 核心变更文件，修改了 forward_extend 方法的分支条件，使 draft_extend_v2 走 decode kernel 路径。

关键符号：forward_extend

## 关键源码片段

### `python/sglang/srt/layers/attention/trtllm_mha_backend.py`

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

```python
# 在 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(...)

```

# 评论区精华

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

- 暂无高价值评论线程

# 风险与影响

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

- 影响：影响范围：仅影响使用 TRT-LLM MHA 后端的 speculative decoding 场景，特别是 draft_extend_v2 阶段。用户无感知，但推理性能可能改善。由于移除了环境变量，运维灵活性略有降低。
- 风险标记：核心路径变更 , 缺少测试覆盖

# 关联脉络

- PR #24863 [Spec v2] Use decode kernel for TRT-LLM MHA draft extend: 上一次尝试相同变更但失败的 PR，本次重新落地。