执行摘要
- 一句话:修正MLX后端单token分块预填充延续误路由为decode
- 推荐动作:建议阅读此PR,以理解状态驱动路由(check decoding_reqs)相比长度启发式的优势,以及通过抽取共享方法同步双路径的设计模式。对于参与MLX后端的开发者,合并后可消除隐式数据损坏风险。
功能与动机
当chunked prefill最后一块只有一个token时,旧逻辑仅通过seq_len>1区分continuation和mixed decode,使得1 token continuation被误判为decode,decode路径不消费输入token,导致真实prompt token被丢弃,生成基于错误的prompt继续。此问题在模型预测与实际token不一致时暴露。详见PR body及repro说明。
实现拆解
-
问题定位:在MlxTpModelWorker中,_forward_batch_generation_mlx和_async_extend_batch都使用seq_len > 1作为判定条件,导致1 token chunk延续进入decode分支。
-
抽取共享路由方法:新增_route_extend_request(self, rid, decoding_rids),根据self._mlx_runner.has_request(rid)和rid in decoding_rids返回prefill、decode或continuation。
-
同步路径修改:在_forward_batch_generation_mlx中构建decoding_rids集合,用路由方法替换原有的长度检查。
-
异步路径修改:在_async_extend_batch中进行相同修改,确保两个路径行为一致。
-
单元测试:新增test_tp_worker_routing.py,使用_FakeRunner mock模型执行器,测试三种路由决策并覆盖同步/异步调用。
关键文件:
test/registered/unit/hardware_backend/mlx/test_tp_worker_routing.py(模块 路由测试;类别 test;类型 test-coverage;符号 _FakeRunner, TestRouting, TestAsyncWiring, test_continuation_one_token_routes_to_extend): 新增的单元测试套件,使用mock runner验证路由决策,覆盖同步和异步路径,是保证修复正确性的关键。
python/sglang/srt/hardware_backend/mlx/tp_worker.py(模块 MLX路由;类别 source;类型 core-logic;符号 _route_extend_request, _forward_batch_generation_mlx, _async_extend_batch): 核心修改文件,新增_route_extend_request方法,修改两个forward路径使用新路由,消除长度启发式。
关键符号:_route_extend_request, _forward_batch_generation_mlx, _async_extend_batch, _FakeRunner.extend, _FakeRunner.decode_batch
关键源码片段
python/sglang/srt/hardware_backend/mlx/tp_worker.py
核心修改文件,新增_route_extend_request方法,修改两个forward路径使用新路由,消除长度启发式。
def _route_extend_request(self, rid: str, decoding_rids: set[str]) -> str:
# 如果 runner 还没有这个 request,则是全新的 prefill
if not self._mlx_runner.has_request(rid):
return "prefill"
# 否则检查是否属于当前 batch 中的真实 decode 步骤(由 scheduler 标记)
if rid in decoding_rids:
return "decode"
# 其他情况都是分块预填充的延续(无论长度是否为 1)
return "continuation"
# 同步路径中的典型使用
decoding_rids = {r.rid for r in (batch.decoding_reqs or [])}
for i, req in enumerate(reqs):
...
route = self._route_extend_request(req.rid, decoding_rids)
if route == "continuation":
next_token = self._mlx_runner.extend(req.rid, req_token_ids, req_new_slots)
extend_rids.append((req.rid, next_token))
elif route == "decode":
decode_rids.append(req.rid)
else: # "prefill"
# 执行 prefill 相关逻辑
...
评论区精华
review中yeahdongcn指出异步路径_async_extend_batch同样存在bug但未被测试固定,建议添加测试或提炼共享helper。noob-se7en采纳建议,将路由决策抽取为_route_extend_request方法,两个路径共同调用,并新增异步路径测试用例,确保两种路径都经过回归验证。
- 建议异步路径测试覆盖或提炼共享helper (testing): 采纳建议:抽取_route_extend_request共享方法,同步和异步路径统一调用,并新增异步路径测试用例覆盖。
风险与影响
- 风险:风险较低:重构提炼共享方法后逻辑简单明确,单元测试覆盖了三种关键场景。但风险包括:
1) batch.decoding_reqs必须由调度器正确维护,否则可能误判;
2) 测试依赖mock runner,未覆盖完整调度器交互;
3) 仅影响MLX后端,不影响其他硬件。
- 影响:
- 用户:Apple Silicon用户在特定prompt长度下使用chunked prefill时,之前可能产生错误推理结果,现在修复。
- 系统:无性能影响,路由开销极小。
- 团队:代码重复消除,可维护性提高;新增测试为后续变更提供安全网。
- 风险标记:MLX后端特有, Mock测试(非端到端), 依赖decoding_reqs正确维护
关联脉络
参与讨论