Prhub

#27475 [spec] Dedup draft `kv_indices` sizing into `spec_utils` helpers

原始 PR 作者 hnyls2002 合并时间 2026-06-07 15:42 文件变更 5 提交数 3 评论 5 代码增减 +82 / -32

执行摘要

提取 draft kv_indices 尺寸计算到 spec_utils 消除重复

将四个 EAGLE draft-decode 后端中重复的 kv_indices 尺寸内联表达式提取为共享辅助函数,使欠分配 bug 类(#27338, #27460)对新后端结构上不可能。

值得精读:展示了如何通过提取共享公式来消除一类隐蔽的缓冲区溢出 bug,是防御性编程的良好实践。未来添加新 EAGLE 后端时应强制使用这两个 helper。

讨论亮点

无 review 评论;作者通过 rerun 测试验证了四个后端的 EAGLE 测试用例(test_spec_eagle_fa3.py、test_spec_eagle_topk.py 等)均通过。

实现拆解

  1. 定义辅助函数:在 python/sglang/srt/speculative/spec_utils.py 中新增 draft_kv_indices_buffer_width(num_seqs, topk, max_context_len)draft_kv_indices_used_len(seq_lens_sum, topk, bs, num_steps),分别计算每步行宽和已用长度。
  2. 更新导入:在四个注意力后端文件(flashinfer_backend.pyflashinfer_mla_backend.pytriton_backend.pyaiter_backend.py)的导入声明中添加 draft_kv_indices_buffer_widthdraft_kv_indices_used_len
  3. 替换内联表达式:在每个后端的 common_template 方法中将容量断言处的 required_kv_indices_len 内联表达式替换为 draft_kv_indices_used_len 调用;将 per-step 切片表达式替换为 draft_kv_indices_used_len 调用;在 init_forward_metadatainit_cuda_graph_state 中将缓冲区维度计算替换为 draft_kv_indices_buffer_width 调用。
  4. 保持行为一致:保留每个站点的 dtype(int32/int64)、device("cuda" vs self.device)以及 torch.empty vs torch.zeros 选择,确保数值结果完全相同。
文件 模块 状态 重要度
python/sglang/srt/speculative/spec_utils.py 工具函数 modified 7.06
python/sglang/srt/layers/attention/flashinfer_backend.py 注意力后端 modified 6.21
python/sglang/srt/layers/attention/flashinfer_mla_backend.py 注意力后端 modified 6.01
python/sglang/srt/layers/attention/triton_backend.py 注意力后端 modified 5.95
python/sglang/srt/layers/attention/aiter_backend.py 注意力后端 modified 5.95

关键符号

draft_kv_indices_buffer_width draft_kv_indices_used_len

关键源码片段

python/sglang/srt/speculative/spec_utils.py core-logic

核心变更文件,添加了两个共享辅助函数,定义了缓冲区宽度和已用长度的计算逻辑。

# 计算 EAGLE draft-decode kv_indices 缓冲区的每步行宽
# 每个分支需要 max_context_len 个 KV 槽,总分支数为 num_seqs * topk
def draft_kv_indices_buffer_width(
    num_seqs: int, topk: int, max_context_len: int
) -> int:
    """Per-step row width of the EAGLE draft-decode kv_indices buffer.    num_seqs * topk branches each attend up to max_context_len KV slots; the topk
    factor is mandatory -- dropping it under-allocates and overflows the row (#27338, #27460).
    """
    return num_seqs * topk * max_context_len
​
​
# 计算 num_steps 步 draft decoding 中实际消耗的 kv_indices 长度
# seq_lens_sum * topk 是初始分支索引数,bs * num_steps 是每步追加的索引数
def draft_kv_indices_used_len(
    seq_lens_sum: int, topk: int, bs: int, num_steps: int
) -> int:
    """kv_indices length used through num_steps draft-decode steps.    bs = topk * num_seqs branches, one index appended per branch per step. Called with
    num_steps = i + 1 (per-step slice) and speculative_num_steps (capacity assert).
    """
    return seq_lens_sum * topk + bs * num_steps
python/sglang/srt/layers/attention/flashinfer_backend.py dependency-wiring

导入新增辅助函数,替换内联表达式,是四个后端中最典型的一个。

# 在 FlashInfer 后端的 common_template 中使用 draft_kv_indices_used_len 进行容量检查
required_kv_indices_len = draft_kv_indices_used_len(
    seq_lens_sum, self.topk, bs, self.speculative_num_steps
)
assert required_kv_indices_len <= kv_indices_buffer.shape[1], (
    f"EAGLE draft kv_indices row too small: need {required_kv_indices_len} "
    f"but row width is {kv_indices_buffer.shape[1]} (topk={self.topk}, "
    f"num_seqs={num_seqs}, seq_lens_sum={seq_lens_sum}, "
    f"num_steps={self.speculative_num_steps}); the buffer must be sized "
    f"max_bs * topk * max_context_len."
)# per-step 切片使用 draft_kv_indices_used_len 计算长度
forward_batch.spec_info.kv_indices = kv_indices_buffer[i][
    : draft_kv_indices_used_len(seq_lens_sum, self.topk, bs, i + 1)
]# 缓冲区分配使用 draft_kv_indices_buffer_width 计算行宽
init_forward_metadata:
    kv_indices_width = draft_kv_indices_buffer_width(
        forward_batch.batch_size, self.topk, self.max_context_len
    )
    kv_indices = torch.empty(
        (self.speculative_num_steps, kv_indices_width),
        dtype=torch.int32,
        device="cuda",
    )init_cuda_graph_state:
    kv_indices_width = draft_kv_indices_buffer_width(
        max_bs, self.topk, self.max_context_len
    )
    self.cuda_graph_kv_indices = torch.zeros(
        (self.speculative_num_steps, kv_indices_width),
        dtype=torch.int32,
        device="cuda",
    )

评论区精华

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

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

风险与影响

低风险:纯重构,数值行为完全等价。但若未来在 spec_utils 中修改这两个 helper,需要同步更新所有调用点(四个后端共 14 处),否则可能导致缓冲区尺寸不匹配。helper 函数的参数顺序和含义需保持文档清晰。

对用户:无用户可见行为变化。对系统:减少代码重复,降低未来引入欠分配 bug 的风险。对团队:新增两个共享 API,新后端可以直接使用,提升开发效率。影响覆盖四个注意力后端,但改动集中且模式一致。

低风险 共享 API 需同步更新

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论