执行摘要
- 一句话:启用 spec 解码下 retraction 排序并清理死代码
- 推荐动作:该 PR 值得精读,尤其适合关注调度器和推测解码的开发者。它展示了如何识别并移除过时的工程限制,同时伴随彻底的死代码清理和测试加固,是良好的 PR 实践。
功能与动机
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 浪费。
实现拆解
- 识别过时约束:在
python/sglang/srt/managers/schedule_batch.py 的 retract_decode 方法中,发现 _get_decode_retraction_order 调用通过 allow_policy_sort 参数禁用了 speculative decoding 下的策略排序,原因是历史限制(filter_batch 只能从尾部过滤)。但分析表明该限制已不成立。
- 移除限制:从
_get_decode_retraction_order 中移除 allow_policy_sort 参数,并删除相关的 if not allow_policy_sort: return sorted_indices 分支,现在始终应用配置的 retraction_policy。
- 清理死代码:从
eagle_info.py、dflash_info_v2.py、ngram_info.py 的 filter_batch 方法中移除 has_been_filtered 参数及其截断分支([:len(new_indices)]),因为该参数仅被 scheduler 调用且始终为 False。
- 删除环境变量:从
python/sglang/srt/environ.py 中删除 SGLANG_SPEC_ENABLE_STRICT_FILTER_CHECK,该变量仅在已删除的截断分支中被读取。
- 补充单元测试:新增
test/registered/unit/managers/test_retraction_order.py,覆盖长度策略、优先级策略以及平局和 None 优先级的情况,验证 _get_decode_retraction_order 的正确性。同时更新 test_dflash_overlap_hostsync.py 以移除已删除的参数。
- 更新文档:从
docs_new/docs/references/environment_variables.mdx 中移除已删除环境变量的文档条目。
关键文件:
test/registered/unit/managers/test_retraction_order.py(模块 单元测试;类别 test;类型 test-coverage;符号 _req, _args, _order, TestRetractionOrder): 新增单元测试,全面覆盖 _get_decode_retraction_order 的排序逻辑(长度策略、优先级策略、平局情况),是验证变更正确性的核心配套。
python/sglang/srt/managers/schedule_batch.py(模块 调度器;类别 source;类型 core-logic;符号 retract_decode, _get_decode_retraction_order, filter_batch): 核心调度逻辑文件:移除 retract_decode 中 spec 下的排序禁用,并清理 _get_decode_retraction_order 的 allow_policy_sort 参数。
python/sglang/srt/speculative/eagle_info.py(模块 推测解码;类别 source;类型 dependency-wiring;符号 filter_batch): 清理 filter_batch 中的 has_been_filtered 参数和截断分支,移除对 envs 的依赖。
python/sglang/srt/environ.py(模块 环境变量;类别 source;类型 configuration;符号 SGLANG_SPEC_ENABLE_STRICT_FILTER_CHECK): 删除环境变量 SGLANG_SPEC_ENABLE_STRICT_FILTER_CHECK,该变量仅被已删除的截断分支读取。
docs_new/docs/references/environment_variables.mdx(模块 文档;类别 docs;类型 documentation): 同步删除环境变量文档条目。
关键符号:_get_decode_retraction_order, retract_decode, filter_batch
关键源码片段
test/registered/unit/managers/test_retraction_order.py
新增单元测试,全面覆盖 _get_decode_retraction_order 的排序逻辑(长度策略、优先级策略、平局情况),是验证变更正确性的核心配套。
import unittest
from types import SimpleNamespace
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_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
核心调度逻辑文件:移除 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
评论区精华
本 PR 未产生大量 review 评论,PR body 自身提供了充分论证:filter_batch 的所有实现都基于任意索引 gather,截断分支已无调用者,因此限制不再成立。
风险与影响
- 风险:风险较低。变更仅影响 speculative decoding 下的 retract 顺序,每个 retract 操作本身不变。单元测试覆盖了排序逻辑的正确性。但若某外部 spec_info 实现依赖旧行为(按到达顺序),可能需要同步调整;仓库内所有实现(eagle_info、dflash_info_v2、ngram_info)均已确认支持任意索引 gather。
- 影响:影响范围限于启用 speculative decoding 的部署。变更后,系统将根据 retraction_policy(默认 length)选择更优的请求进行 retract,优先 retract 输出较短的请求,从而减少 token 浪费。非 speculative 路径无变化。
- 风险标记:核心路径变更, spec 排序行为变更, 环境变量删除, 依赖清理
关联脉络
- PR #27977 [Speculative Decoding] Remove v1 scheduler paths: PR body 提及此 PR 移除了 spec v1 调度路径,导致 has_been_filtered=True 分支不再有调用者。
- PR #3988 [Speculative Decoding] Initial spec v1: PR body 追溯约束起源,spec v1 中 filter_batch 仅截断,因此只能从尾部 retract。
- PR #9252 [Speculative Decoding] Add gather branch to filter_batch: PR body 提及此 PR 添加了 gather 分支(has_been_filtered=False),使排序成为可能。
- PR #30573 [Scheduler] Extract decode retraction order helper: PR body 提及此 PR 抽取了 _get_decode_retraction_order 辅助方法,但保留了约束条件作为化石。
参与讨论