Prhub

#30181 [MLX] Fix single-token chunked-prefill continuation misrouted as decode

原始 PR 作者 noob-se7en 合并时间 2026-07-07 11:31 文件变更 2 提交数 3 评论 4 代码增减 +293 / -26

执行摘要

修正 MLX 后端单 token 分块预填充延续误路由为 decode

当chunked prefill最后一块只有一个token时,旧逻辑仅通过seq_len>1区分continuation和mixed decode,使得1 token continuation被误判为decode,decode路径不消费输入token,导致真实prompt token被丢弃,生成基于错误的prompt继续。此问题在模型预测与实际token不一致时暴露。详见PR body及repro说明。

建议阅读此PR,以理解状态驱动路由(check decoding_reqs)相比长度启发式的优势,以及通过抽取共享方法同步双路径的设计模式。对于参与MLX后端的开发者,合并后可消除隐式数据损坏风险。

讨论亮点

review中yeahdongcn指出异步路径_async_extend_batch同样存在bug但未被测试固定,建议添加测试或提炼共享helper。noob-se7en采纳建议,将路由决策抽取为_route_extend_request方法,两个路径共同调用,并新增异步路径测试用例,确保两种路径都经过回归验证。

实现拆解

  1. 问题定位:在MlxTpModelWorker中,_forward_batch_generation_mlx_async_extend_batch都使用seq_len > 1作为判定条件,导致1 token chunk延续进入decode分支。

  2. 抽取共享路由方法:新增_route_extend_request(self, rid, decoding_rids),根据self._mlx_runner.has_request(rid)rid in decoding_rids返回prefilldecodecontinuation

  3. 同步路径修改:在_forward_batch_generation_mlx中构建decoding_rids集合,用路由方法替换原有的长度检查。

  4. 异步路径修改:在_async_extend_batch中进行相同修改,确保两个路径行为一致。

  5. 单元测试:新增test_tp_worker_routing.py,使用_FakeRunner mock模型执行器,测试三种路由决策并覆盖同步/异步调用。

文件 模块 状态 重要度
test/registered/unit/hardware_backend/mlx/test_tp_worker_routing.py 路由测试 added 7.9
python/sglang/srt/hardware_backend/mlx/tp_worker.py MLX 路由 modified 7.55

关键符号

_route_extend_request _forward_batch_generation_mlx _async_extend_batch _FakeRunner.extend _FakeRunner.decode_batch

关键源码片段

python/sglang/srt/hardware_backend/mlx/tp_worker.py core-logic

核心修改文件,新增 _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 相关逻辑
        ...

评论区精华

建议异步路径测试覆盖或提炼共享 helper 测试

yeahdongcn 指出异步路径 _async_extend_batch 同样存在 bug 但未被测试固定,建议添加测试或提炼共享 helper 确保一致。

结论:采纳建议:抽取 _route_extend_request 共享方法,同步和异步路径统一调用,并新增异步路径测试用例覆盖。 · 已解决

风险与影响

风险较低:重构提炼共享方法后逻辑简单明确,单元测试覆盖了三种关键场景。但风险包括:

1) batch.decoding_reqs必须由调度器正确维护,否则可能误判;
2) 测试依赖mock runner,未覆盖完整调度器交互;
3) 仅影响MLX后端,不影响其他硬件。

  • 用户:Apple Silicon用户在特定prompt长度下使用chunked prefill时,之前可能产生错误推理结果,现在修复。
  • 系统:无性能影响,路由开销极小。
  • 团队:代码重复消除,可维护性提高;新增测试为后续变更提供安全网。
MLX 后端特有 Mock 测试(非端到端) 依赖 decoding_reqs 正确维护

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论