diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 0ddccdf46..7f2de88cb 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -129,15 +129,17 @@ class DecodeReqToTokenPool: return len(self.free_slots) def alloc(self, reqs: List["Req"]) -> Optional[List[int]]: - chunked = [i for i, r in enumerate(reqs) if r.req_pool_idx is not None] + # Indices of reqs that already have a req_pool_idx and will reuse + # their existing slot (e.g. chunked prefill continuing across chunks). + reusing = [i for i, r in enumerate(reqs) if r.req_pool_idx is not None] assert ( - len(chunked) <= 1 + len(reusing) <= 1 ), "only one chunked request may reuse req_pool_idx in a batch" assert all( - reqs[i].is_chunked > 0 or reqs[i].kv_committed_len > 0 for i in chunked - ), "request has req_pool_idx but is not chunked" + reqs[i].is_chunked > 0 or reqs[i].kv_committed_len > 0 for i in reusing + ), "reusing request must be chunked or have committed KV" - need_size = len(reqs) - len(chunked) + need_size = len(reqs) - len(reusing) if need_size > len(self.free_slots): return None select_index = self.free_slots[:need_size] diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 10b9d9b49..5312b2edb 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -153,16 +153,18 @@ class ReqToTokenPool: return len(self.free_slots) def alloc(self, reqs: list[Req]) -> Optional[List[int]]: - chunked = [i for i, r in enumerate(reqs) if r.req_pool_idx is not None] + # Indices of reqs that already have a req_pool_idx and will reuse + # their existing slot (e.g. chunked prefill continuing across chunks). + reusing = [i for i, r in enumerate(reqs) if r.req_pool_idx is not None] if not any(r.is_dllm() for r in reqs): assert ( - len(chunked) <= 1 + len(reusing) <= 1 ), "only one chunked request may reuse req_pool_idx in a batch" assert all( - reqs[i].is_chunked > 0 or reqs[i].kv_committed_len > 0 for i in chunked - ), "request has req_pool_idx but is not chunked" + reqs[i].is_chunked > 0 or reqs[i].kv_committed_len > 0 for i in reusing + ), "reusing request must be chunked or have committed KV" - need_size = len(reqs) - len(chunked) + need_size = len(reqs) - len(reusing) if need_size > len(self.free_slots): return None select_index = self.free_slots[:need_size]