Prhub

#27799 [Spec] `NGRAMWorker` on `BaseSpecWorker`; algo-owned verify-tree shape params

原始 PR 作者 hnyls2002 合并时间 2026-06-11 03:38 文件变更 6 提交数 5 评论 5 代码增减 +54 / -19

执行摘要

重构 NGRAMWorker,统一推测解码树形参数接口

PR #27799 旨在将上一步 #17260 引入的 NGRAM 推测解码实现接入标准接口,并使其树形参数(验证树深度、分支因子)通过算法类型的属性暴露,从而让共享的 sample() 和 KV 预分配逻辑不再需要 is_ngram() 分支。这样不仅减少了代码死角,也为未来引入 dflash、Medusa 等树形算法铺平了道路。

  • 优先修复 critical 问题:common.py 中对 server_args.speculative_algorithm 做 None 检查,或给 SpeculativeAlgorithm.from_string 增加 None 安全处理。
  • 值得学习的模式: 将树形参数提升到 verify-input 抽象上,通过多态消除条件分支,是一个很好的接口设计范例,适合在其他需要算法多态的场景中借鉴。
  • 补充测试:SpeculativeAlgorithm.has_draft_kv() 和 verify-input 属性编写简单单元测试,固化接口契约。
讨论亮点

gemini-code-assist[bot] 提出了三个评审意见,均未在合并前采纳:

  • Critical — from_string(None) 崩溃风险:server_args.speculative_algorithm 为 None 时,get_alloc_len_per_decode 中的 SpeculativeAlgorithm.from_string(None) 会抛出 AttributeError,导致任何非推测解码请求都崩溃。建议调用前判断 None。
  • Medium — GPU 循环创建 tensor 的性能问题:ngram_worker.py 的一次循环中,逐请求创建 GPU tensor 并 torch.cat,可能引入额外延迟,建议改用 CPU numpy 构建后统一传输。
  • Medium — 除零错误风险:speculative_hook.py 的 NGRAM 初始化中,speculative_num_draft_tokens // speculative_eagle_topk 可能因 topk 为 0 或 None 而异常,建议添加安全 fallback。

实现拆解

步骤 1:让 NGRAMWorker 继承 BaseSpecWorker

ngram_worker.py 中,类声明改为 class NGRAMWorker(BaseSpecWorker),新增 target_worker 属性和 draft_worker 属性,前者返回底层目标 worker,后者返回 None 表示无 draft 模型。同时调整导入语句,引入 BaseDraftWorkerBaseSpecWorker

步骤 2:将验证树形状参数改为多态属性

eagle_info_v2.pyEagleVerifyInputV2Mixin 中新增 max_tree_depthtree_topk 属性,默认值由 spec_steps 和 topk 计算;sample() 方法中直接使用 self.max_tree_depthself.tree_topk 取代原先的 is_ngram() 条件判断。

ngram_info.pyNgramVerifyInput 中覆盖这两个属性:max_tree_depth 返回 draft_token_num(NGRAM 树受节点预算限制,单次匹配可能链满所有 token),tree_topk 返回 -1(无固定分支因子)。

步骤 3:引入算法级别的 has_draft_kv() 方法

spec_info.pySpeculativeAlgorithm 枚举和 spec_registry.py 的基类中添加 has_draft_kv() 抽象。NGRAM 算法返回 False(其 draft 树只存在于验证掩码中,不占用实际 KV 缓存),其他算法返回 True。该值在 mem_cache/common.pyget_alloc_len_per_decode() 中使用,替代原有的字符串比较 is_ngram,决定是否启用 per-topk 页面复制预留。

步骤 4:简化 KV 预分配逻辑

common.py 中,将原来基于 server_args.speculative_algorithm.upper() == "NGRAM" 的条件改为使用 spec_algo.has_draft_kv(),使逻辑更健壮且算法无关。

步骤 5:无测试新增,但现有套件覆盖正常路径

本次 PR 未添加单元测试,但 test_spec_ngram.pytest_spec_eagle_parity.py 等集成测试在 CI 中通过,保证了基本功能不退化。

文件 模块 状态 重要度
python/sglang/srt/speculative/eagle_info_v2.py 推测解码 modified 7.29
python/sglang/srt/speculative/ngram_worker.py 推测解码 modified 7.29
python/sglang/srt/speculative/ngram_info.py 推测解码 modified 6.38
python/sglang/srt/mem_cache/common.py 推测解码 modified 6.07
python/sglang/srt/speculative/spec_info.py 推测解码 modified 5.69
python/sglang/srt/speculative/spec_registry.py 推测解码 modified 5.49

关键符号

NGRAMWorker.__init__ NGRAMWorker.target_worker NGRAMWorker.draft_worker EagleVerifyInputV2Mixin.max_tree_depth EagleVerifyInputV2Mixin.tree_topk NgramVerifyInput.max_tree_depth NgramVerifyInput.tree_topk SpeculativeAlgorithm.has_draft_kv SpeculativeAlgorithmRegistry.has_draft_kv

关键源码片段

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

核心变更文件:新增 max_tree_depth/tree_topk 属性到 EagleVerifyInputV2Mixin,消除 sample() 中对 is_ngram() 的依赖

@dataclass
class EagleVerifyInputV2Mixin:
    @property
    def max_tree_depth(self: EagleVerifyInput) -> int:
        """最长根到叶的验证树链,包含根节点;用于确定 accept_index 行宽度。
        EAGLE 树受 draft 循环深度限制,spec_steps + 1。"""
        return self.spec_steps + 1
​
    @property
    def tree_topk(self: EagleVerifyInput) -> int:
        """传递给树验证核的分支因子;-1 表示不规则树。"""
        return self.topk
​
    def sample(self, ...):
        # ...
        candidates = self.draft_token.reshape(bs, self.draft_token_num)
        predict_shape = list(next_token_logits.shape)[:-1]
        predict = torch.zeros(predict_shape, dtype=torch.int32, device=device).flatten()
        # 使用 self.max_tree_depth 替代条件分支
        accept_index = torch.full(
            (bs, self.max_tree_depth), -1, dtype=torch.int32, device=device
        )
        # ...
        if tree_mask_method == "custom_mask_tree":
            # ...
            topk = self.tree_topk
            # ...
        # ...
python/sglang/srt/speculative/ngram_worker.py core-logic

让 NGRAMWorker 继承 BaseSpecWorker,并实现 target_worker/draft_worker 接口属性

class NGRAMWorker(BaseSpecWorker):
    def __init__(self, server_args, gpu_id, tp_rank, dp_rank,
                 moe_ep_rank, attn_cp_rank, moe_dp_rank, nccl_port,
                 target_worker):
        # ...
        self._target_worker = target_worker # 改为 private 属性
        # ...
​
    @property
    def target_worker(self) -> TpModelWorker:
        return self._target_worker
​
    @property
    def draft_worker(self) -> Optional[BaseDraftWorker]:
        # NGRAM 没有 draft 模型;draft 来自 CPU 端语料库
        return None
python/sglang/srt/speculative/ngram_info.py core-logic

覆盖 max_tree_depth/tree_topk 属性以提供 NGRAM 特有的树形状

@dataclass
class NgramVerifyInput(SpecVerInput):
    # ...
​
    @property
    def max_tree_depth(self) -> int:
        # NGRAM 树受节点预算限制没有深度上限:语料库 BFS 只会在节点预算
        # 用尽时停止,因此单次长匹配可以链接所有 draft_token_num 个节点
        # (spec_steps 对此树无意义)。
        return self.draft_token_num
​
    @property
    def tree_topk(self) -> int:
        # 不规则树:每层分支跟随语料库匹配结果。
        return -1

评论区精华

未保护的 from_string 调用导致正常解码崩溃 正确性

gemini-code-assist[bot] 指出:当 server_args.speculative_algorithm 为 None 时,调用 SpeculativeAlgorithm.from_string(None) 会抛出 AttributeError/ValueError,导致无推测解码的请求完全崩溃。建议加 None 判断。

结论:作者未在 PR 中回复或修复该问题,合并后代码仍带有此 bug。 · 待处理

GPU tensor 循环创建性能问题 性能

gemini-code-assist[bot] 指出:在 ngram_worker.py 的 batch 循环中使用 torch.ones 和 torch.cat 创建 GPU tensor 会产生多次内核启动开销,建议用 CPU numpy 构建后一次性传输。

结论:未采纳,作者未回应。 · 待处理

除零错误风险 正确性

在 speculative_hook.py 中,当 server_args.speculative_eagle_topk 为 None 或 0 时,除法 speculative_num_draft_tokens // speculative_eagle_topk 可能抛出 TypeError 或 ZeroDivisionError。建议加安全 fallback。

结论:未修复,合并后代码仍存在此风险。 · 待处理

风险与影响

核心路径崩溃风险(Critical): 如 review 指出的 from_string(None) 调用未加保护,一旦用户不使用推测解码,每个 decode 步都会触发异常,完全阻塞正常推理。该项目目前已合并该 PR,此问题仍存在,属于严重回归。

缺少测试覆盖: 本次重构涉及接口契约变化,但未针对新属性或 has_draft_kv() 编写单元测试,回归依赖已有的 test_spec_ngram.py 等,若未来有其他算法接入可能漏测。

潜在性能回退: review 指出 ngram_worker.py 中的 batch 循环内创建 GPU tensor,可能在大 batch 场景下引入不可忽视的开销,但官方未做 benchmark。

内部 API 影响: NGRAMWorker 现在必须遵守 BaseSpecWorker 接口,若外部有直接使用 NGRAMWorker.target_worker(现在改为 _target_worker 后通过 property 暴露)的代码,需要进行适配。但调度器已通过接口使用,影响可控。

用户影响: 无直接可见变更,推测解码行为不变。

系统影响: 新的多态抽象使添加新算法(如 dflash、Medusa)时只需提供对应 verify-input 类型并实现 max_tree_depth/tree_topk/has_draft_kv,无需修改共享 sample() 和内存分配逻辑,降低耦合。

核心路径崩溃风险 缺少测试覆盖 GPU 性能回退可能

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论