Prhub

#32023 [Scheduler] Enable decode retraction ordering under speculative decoding

原始 PR 作者 hnyls2002 合并时间 2026-07-23 15:42 文件变更 8 提交数 4 评论 4 代码增减 +71 / -57

执行摘要

启用 spec 解码下 retraction 排序并清理死代码

PR body 指出,retract_decode 中禁用了 speculative decoding 下的 retraction_order 策略,原因是 filter_batch 只能从尾部过滤请求。但 scheduler 传递 has_been_filtered=False 后,所有 spec_info 实现都使用 gather 索引,没有头部限制。此约束来自旧 spec v1 路径,现已过期。移除限制可让 spec 解码下也按配置策略(默认 shortest-output first)进行 retract,减少 token 浪费。

该 PR 值得精读,尤其适合关注调度器和推测解码的开发者。它展示了如何识别并移除过时的工程限制,同时伴随彻底的死代码清理和测试加固,是良好的 PR 实践。

讨论亮点

本 PR 未产生大量 review 评论,PR body 自身提供了充分论证:filter_batch 的所有实现都基于任意索引 gather,截断分支已无调用者,因此限制不再成立。

实现拆解

  1. 识别过时约束:在 python/sglang/srt/managers/schedule_batch.pyretract_decode 方法中,发现 _get_decode_retraction_order 调用通过 allow_policy_sort 参数禁用了 speculative decoding 下的策略排序,原因是历史限制(filter_batch 只能从尾部过滤)。但分析表明该限制已不成立。
  2. 移除限制:从 _get_decode_retraction_order 中移除 allow_policy_sort 参数,并删除相关的 if not allow_policy_sort: return sorted_indices 分支,现在始终应用配置的 retraction_policy。
  3. 清理死代码:从 eagle_info.pydflash_info_v2.pyngram_info.pyfilter_batch 方法中移除 has_been_filtered 参数及其截断分支([:len(new_indices)]),因为该参数仅被 scheduler 调用且始终为 False
  4. 删除环境变量:从 python/sglang/srt/environ.py 中删除 SGLANG_SPEC_ENABLE_STRICT_FILTER_CHECK,该变量仅在已删除的截断分支中被读取。
  5. 补充单元测试:新增 test/registered/unit/managers/test_retraction_order.py,覆盖长度策略、优先级策略以及平局和 None 优先级的情况,验证 _get_decode_retraction_order 的正确性。同时更新 test_dflash_overlap_hostsync.py 以移除已删除的参数。
  6. 更新文档:从 docs_new/docs/references/environment_variables.mdx 中移除已删除环境变量的文档条目。
文件 模块 状态 重要度
test/registered/unit/managers/test_retraction_order.py 单元测试 added 7.27
python/sglang/srt/managers/schedule_batch.py 调度器 modified 6.06
python/sglang/srt/speculative/eagle_info.py 推测解码 modified 6.55
python/sglang/srt/environ.py 环境变量 modified 4.35
docs_new/docs/references/environment_variables.mdx 文档 modified 2.38

关键符号

_get_decode_retraction_order retract_decode filter_batch

关键源码片段

test/registered/unit/managers/test_retraction_order.py test-coverage

新增单元测试,全面覆盖 _get_decode_retraction_order 的排序逻辑(长度策略、优先级策略、平局情况),是验证变更正确性的核心配套。

import unittest
from types import SimpleNamespacefrom sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCaseregister_cpu_ci(est_time=2, suite="base-a-test-cpu")
​
​
# 辅助函数:创建模拟请求。注意 output_len 影响排序,input_len 作为平局时的次要键
def _req(output_len: int, input_len: int = 8, priority=None):
    return SimpleNamespace(
        output_ids=[0] * output_len,
        origin_input_ids=[0] * input_len,
        priority=priority,
    )
​
​
# 辅助函数:创建模拟 server_args
def _args(policy: str = "length", low_first: bool = False):
    return SimpleNamespace(
        retraction_policy=policy,
        schedule_low_priority_values_first=low_first,
    )
​
​
# 直接调用静态方法 _get_decode_retraction_order
def _order(reqs, args):
    return ScheduleBatch._get_decode_retraction_order(reqs, args)
​
​
class TestRetractionOrder(CustomTestCase):
    """注意 retraction 循环从列表末尾 pop,因此返回列表的最后一个元素是最先被 retract 的请求。"""
​
    def test_length_policy_retracts_shortest_output_first(self):
        # 按输出长度升序排序,所以索引 1(output_len=1)应排在最后(最先被 retract)
        reqs = [_req(5), _req(1), _req(3)]
        self.assertEqual(_order(reqs, _args()), [0, 2, 1])
​
    def test_length_policy_tie_breaks_on_longer_input(self):
        # 输出长度相同(4),则输入更长的请求被优先 retract(释放更多 token)
        reqs = [_req(4, input_len=10), _req(4, input_len=20)]
        self.assertEqual(_order(reqs, _args()), [0, 1])
​
    def test_priority_policy_low_values_first(self):
        # 低优先级值优先(low_first=True);None 视为最不重要
        reqs = [_req(4, priority=2), _req(4, priority=0), _req(4, priority=None)]
        self.assertEqual(_order(reqs, _args("priority", low_first=True)), [1, 0, 2])
​
    def test_priority_policy_high_values_first(self):
        # 高优先级值优先(low_first=False)
        reqs = [_req(4, priority=2), _req(4, priority=0), _req(4, priority=None)]
        self.assertEqual(_order(reqs, _args("priority", low_first=False)), [0, 1, 2])
​
    def test_priority_ties_fall_back_to_length(self):
        # 优先级相同时,退回到长度策略
        reqs = [_req(1, priority=1), _req(5, priority=1)]
        self.assertEqual(_order(reqs, _args("priority", low_first=True)), [1, 0])
​
​
if __name__ == "__main__":
    unittest.main()
python/sglang/srt/managers/schedule_batch.py core-logic

核心调度逻辑文件:移除 retract_decode 中 spec 下的排序禁用,并清理 _get_decode_retraction_order 的 allow_policy_sort 参数。

def retract_decode(
    self, server_args: ServerArgs
) -> Tuple[List[Req], float, List[Req]]:
    """当 decode 内存不足时 retract 请求。"""
    # 现在始终调用 _get_decode_retraction_order,不再根据 spec 禁用排序
    sorted_indices = self._get_decode_retraction_order(self.reqs, server_args)
​
    retracted_reqs = []
    first_iter = True
    while first_iter or (
        not self.check_decode_mem(selected_indices=sorted_indices)
    ):
        if len(sorted_indices) == 1:
            # 始终保留至少一个请求
            break
        first_iter = False
        idx = sorted_indices.pop() # 从末尾取,即策略中“最不偏好”的请求
        req = self.reqs[idx]
        retracted_reqs.append(req)
        self.release_req(idx, len(sorted_indices), server_args)
​
    # ... 后续 abort 逻辑和 filter_batch 调用不变 ...@staticmethod
def _get_decode_retraction_order(
    reqs: List[Req], server_args: ServerArgs
) -> List[int]:
    """返回索引列表,从最偏好保留到最不偏好保留。
    循环从末尾 pop,因此最不偏好的请求被最先 retract。"""
    sorted_indices = list(range(len(reqs)))
​
    # 长度策略的键函数:优先保留输出较长的请求(更不浪费),平局时保留输入较短的
    def length_key(req: Req) -> Tuple[int, int]:
        return (len(req.output_ids), -len(req.origin_input_ids))
​
    if server_args.retraction_policy == "priority":
        priority_sign = 1 if server_args.schedule_low_priority_values_first else -1
​
        def retraction_key(req: Req) -> Tuple[int, int, int]:
            priority = req.priority
            if priority is None:
                priority = (
                    sys.maxsize
                    if server_args.schedule_low_priority_values_first
                    else -sys.maxsize - 1
                )
            return (priority * (-priority_sign), *length_key(req))
​
        sorted_indices.sort(
            key=lambda i: retraction_key(reqs[i]),
            reverse=True,
        )
        return sorted_indices
​
    # 默认 length 策略:按输出长度降序(输出短的优先站在列表尾部)
    sorted_indices.sort(
        key=lambda i: length_key(reqs[i]),
        reverse=True,
    )
    return sorted_indices

评论区精华

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

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

风险与影响

风险较低。变更仅影响 speculative decoding 下的 retract 顺序,每个 retract 操作本身不变。单元测试覆盖了排序逻辑的正确性。但若某外部 spec_info 实现依赖旧行为(按到达顺序),可能需要同步调整;仓库内所有实现(eagle_info、dflash_info_v2、ngram_info)均已确认支持任意索引 gather。

影响范围限于启用 speculative decoding 的部署。变更后,系统将根据 retraction_policy(默认 length)选择更优的请求进行 retract,优先 retract 输出较短的请求,从而减少 token 浪费。非 speculative 路径无变化。

核心路径变更 spec 排序行为变更 环境变量删除 依赖清理

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论