diff --git a/python/sglang/benchmark/one_batch.py b/python/sglang/benchmark/one_batch.py index bdd2d9b85..f21ed2faf 100644 --- a/python/sglang/benchmark/one_batch.py +++ b/python/sglang/benchmark/one_batch.py @@ -432,9 +432,9 @@ def prepare_extend_inputs_for_correctness_test( req: Req = reqs[i] req.full_untruncated_fill_ids.extend(input_ids[i][bench_args.cut_len :]) if model_runner is not None: - # Use req.req_pool_idx instead of i to handle slot 0 padding correctly + # Use req.kv.req_pool_idx instead of i to handle slot 0 padding correctly req.prefix_indices = model_runner.req_to_token_pool.req_to_token[ - req.req_pool_idx, : bench_args.cut_len + req.kv.req_pool_idx, : bench_args.cut_len ].to(req.prefix_indices.dtype) req.logprob_start_len = -1 req.set_extend_range( diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index d75239116..ed5f836e2 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -188,14 +188,10 @@ class DecodeReqToTokenPool: def alloc(self, reqs: List[Req]) -> Optional[List[int]]: # 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.kv.req_pool_idx is not None] - assert ( - len(reusing) <= 1 - ), "only one chunked request may reuse req_pool_idx in a batch" + reusing = [i for i, r in enumerate(reqs) if r.kv.holds_kv] assert all( - reqs[i].inflight_middle_chunks > 0 or reqs[i].kv.kv_committed_len > 0 - for i in reusing - ), "reusing request must be chunked or have committed KV" + reqs[i].kv.kv_allocated_len > 0 for i in reusing + ), "a reused row must carry allocated KV" need_size = len(reqs) - len(reusing) if need_size > len(self.free_slots): @@ -204,14 +200,14 @@ class DecodeReqToTokenPool: self.free_slots = self.free_slots[need_size:] offset = 0 for r in reqs: - if r.kv.req_pool_idx is None: + if not r.kv.holds_kv: r.kv.req_pool_idx = select_index[offset] self.req_generation[r.kv.req_pool_idx] += 1 offset += 1 return [r.kv.req_pool_idx for r in reqs] def free(self, req: Req): - assert req.kv.req_pool_idx is not None, "request must have req_pool_idx" + assert req.kv.holds_kv, "request must have req_pool_idx" self.free_slots.append(req.kv.req_pool_idx) req.kv.req_pool_idx = None diff --git a/python/sglang/srt/hardware_backend/npu/dsv4/dsv4_req_to_token_pool.py b/python/sglang/srt/hardware_backend/npu/dsv4/dsv4_req_to_token_pool.py index 00b40f081..86b3a8913 100644 --- a/python/sglang/srt/hardware_backend/npu/dsv4/dsv4_req_to_token_pool.py +++ b/python/sglang/srt/hardware_backend/npu/dsv4/dsv4_req_to_token_pool.py @@ -83,7 +83,7 @@ class DSV4ReqToTokenTablesMixin: Prefix matching can happen before a request slot is allocated, so the page ids are temporarily carried by ``Req`` and installed by ``alloc``. """ - if req.kv.req_pool_idx is None: + if not req.kv.holds_kv: req.c128_prefix_page_ids = page_ids return self._dsv4_allocator.replace_req_c128_prefix( @@ -91,7 +91,7 @@ class DSV4ReqToTokenTablesMixin: ) def alloc(self, reqs): - fresh = [req.kv.req_pool_idx is None for req in reqs] + fresh = [not req.kv.holds_kv for req in reqs] indices = super().alloc(reqs) if indices is None: return None diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 058854e23..ecbc6dd6d 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -3124,7 +3124,7 @@ class Scheduler( if self.chunked_req is not req: # Already past chunked prefill; the running-batch abort path handles # it. Drop the marker once the request is actually gone. - if req.finished() or req.kv.req_pool_idx is None: + if req.finished() or not req.kv.holds_kv: self._pending_chunked_abort_req = None return diff --git a/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py b/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py index f653c5f68..5529b770f 100644 --- a/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py +++ b/python/sglang/srt/managers/scheduler_components/pool_stats_observer.py @@ -172,7 +172,7 @@ class SchedulerPoolStatsObserver: if batch is None or batch.is_empty(): continue for req in batch.reqs: - if req.kv.req_pool_idx is not None: + if req.kv.holds_kv: idxs.add(req.kv.req_pool_idx) return idxs diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index 51f8b34c4..009dd2ce2 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -31,6 +31,7 @@ from sglang.srt.managers.utils import ( get_logprob_dict_from_result, get_logprob_from_pp_outputs, ) +from sglang.srt.mem_cache.common import release_kv_cache from sglang.srt.model_executor.forward_batch_info import ( ForwardBatch, ForwardMode, @@ -724,14 +725,7 @@ class SchedulerPPMixin: # Release KV and Mamba cache if req.kv.holds_kv: - kv_indices = self.req_to_token_pool.req_to_token[ - req.kv.req_pool_idx, : req.extend_range.end - ] - self.token_to_kv_pool_allocator.free(kv_indices) - if req.kv.holds_mamba: - self.req_to_token_pool.free_mamba_cache(req) - self.req_to_token_pool.free(req) - req.kv.mark_kv_released() + release_kv_cache(req, self.tree_cache, is_insert=False) logger.info( f"[PP Dynamic Chunk] [PP0] Profiled {len(seq_lens)} samples: " diff --git a/python/sglang/srt/mem_cache/allocation.py b/python/sglang/srt/mem_cache/allocation.py index fe8a2e858..5e109a4c8 100644 --- a/python/sglang/srt/mem_cache/allocation.py +++ b/python/sglang/srt/mem_cache/allocation.py @@ -296,10 +296,7 @@ def alloc_for_extend( reuse_kv = None if batch.is_dllm(): - reuse_kv = [ - r.kv.req_pool_idx is not None and bool(r.dllm_incomplete_ids) - for r in batch.reqs - ] + reuse_kv = [r.kv.holds_kv and bool(r.dllm_incomplete_ids) for r in batch.reqs] # Create tensors for allocation pin_memory = is_pin_memory_available(batch.device) diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 9e9ede5cc..f11b9856d 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -292,24 +292,17 @@ class ReqToTokenPool: def alloc(self, reqs: list[Req]) -> Optional[List[int]]: # 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.kv.req_pool_idx is not None] - # NOTE: this check is relaxed temporarily - # https://github.com/sgl-project/sglang/pull/20476 - # if not any(r.is_dllm() for r in reqs): - # assert ( - # sum(1 for i in reusing if reqs[i].inflight_middle_chunks > 0) <= 1 - # ), "only one chunked request may reuse req_pool_idx in a batch" + reusing = [i for i, r in enumerate(reqs) if r.kv.holds_kv] assert all( - reqs[i].inflight_middle_chunks > 0 or reqs[i].kv.kv_committed_len > 0 - for i in reusing - ), "reusing request must be chunked or have committed KV" + reqs[i].kv.kv_allocated_len > 0 for i in reusing + ), "a reused row must carry allocated KV" select_index = self.alloc_rows(len(reqs) - len(reusing)) if select_index is None: return None offset = 0 for r in reqs: - if r.kv.req_pool_idx is None: + if not r.kv.holds_kv: r.kv.req_pool_idx = select_index[offset] offset += 1 return [r.kv.req_pool_idx for r in reqs] @@ -338,7 +331,7 @@ class ReqToTokenPool: self.free_slots.extend(indices) def free(self, req: Req): - assert req.kv.req_pool_idx is not None, "request must have req_pool_idx" + assert req.kv.holds_kv, "request must have req_pool_idx" self.free_rows([req.kv.req_pool_idx]) req.kv.req_pool_idx = None diff --git a/python/sglang/srt/mem_cache/radix_cache_cpp.py b/python/sglang/srt/mem_cache/radix_cache_cpp.py index 547214104..52bdd9013 100644 --- a/python/sglang/srt/mem_cache/radix_cache_cpp.py +++ b/python/sglang/srt/mem_cache/radix_cache_cpp.py @@ -185,7 +185,7 @@ class RadixCacheCpp(BasePrefixCache): ): """Cache request when it finishes.""" self._reject_cache_salt(req.cache_salt) - assert req.kv.req_pool_idx is not None + assert req.kv.holds_kv token_ids = (req.origin_input_ids + req.output_ids)[:kv_len_to_handle] kv_indices = self.req_to_token_pool.req_to_token[ req.kv.req_pool_idx, :kv_len_to_handle @@ -223,7 +223,7 @@ class RadixCacheCpp(BasePrefixCache): def cache_unfinished_req(self, req: Req, chunked=False): """Cache request when it is unfinished.""" self._reject_cache_salt(req.cache_salt) - assert req.kv.req_pool_idx is not None + assert req.kv.holds_kv token_ids = req.get_fill_ids() prefill_len = len(token_ids) # prefill only (maybe chunked) kv_indices = self.req_to_token_pool.req_to_token[ diff --git a/python/sglang/srt/mem_cache/sparsity/core/sparse_coordinator.py b/python/sglang/srt/mem_cache/sparsity/core/sparse_coordinator.py index cb5d91755..993b4a16b 100644 --- a/python/sglang/srt/mem_cache/sparsity/core/sparse_coordinator.py +++ b/python/sglang/srt/mem_cache/sparsity/core/sparse_coordinator.py @@ -142,7 +142,7 @@ class SparseCoordinator: Registers the request in the state tracker to enable sparse attention processing. """ - if req.kv.req_pool_idx is not None: + if req.kv.holds_kv: self.states.register(req.kv.req_pool_idx, len(req.origin_input_ids)) def on_request_end(self, req: "Req") -> None: @@ -150,7 +150,7 @@ class SparseCoordinator: Handle request end event. Called when a request is completed or aborted. Cleans up request-specific state and releases resources. """ - if req.kv.req_pool_idx is None: + if not req.kv.holds_kv: return self.states.clear(req.kv.req_pool_idx) diff --git a/test/registered/unit/layers/test_minicpm_sparse_cache.py b/test/registered/unit/layers/test_minicpm_sparse_cache.py index 2690ad7a1..3dffb3729 100644 --- a/test/registered/unit/layers/test_minicpm_sparse_cache.py +++ b/test/registered/unit/layers/test_minicpm_sparse_cache.py @@ -68,7 +68,7 @@ def make_pool_and_req(capacity: int = 64): ) req = SimpleNamespace( inflight_middle_chunks=0, - kv=SimpleNamespace(req_pool_idx=None), + kv=ReqKvInfo(), ) req_pool_idx = pool.alloc([req])[0] return pool, req, req_pool_idx, allocator diff --git a/test/registered/unit/mem_cache/test_decode_retraction_backup.py b/test/registered/unit/mem_cache/test_decode_retraction_backup.py index 24e53ffee..4b9f456c2 100644 --- a/test/registered/unit/mem_cache/test_decode_retraction_backup.py +++ b/test/registered/unit/mem_cache/test_decode_retraction_backup.py @@ -3,6 +3,7 @@ from types import SimpleNamespace import torch +from sglang.srt.managers.schedule_batch import ReqKvInfo from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator from sglang.srt.mem_cache.cache_init_params import CacheInitParams from sglang.srt.mem_cache.common import retraction_backup @@ -129,9 +130,7 @@ class TestDecodeRetractionBackup(unittest.TestCase): ) def _admit_req(self, env, num_tokens: int): - req = SimpleNamespace( - rid="request", kv=SimpleNamespace(req_pool_idx=None), seqlen=num_tokens + 1 - ) + req = SimpleNamespace(rid="request", kv=ReqKvInfo(), seqlen=num_tokens + 1) self.assertIsNotNone(env.req_to_token_pool.alloc([req])) source_indices = env.allocator.alloc(num_tokens) self.assertIsNotNone(source_indices) diff --git a/test/registered/unit/mem_cache/test_dllm_fdfo_kv_reuse.py b/test/registered/unit/mem_cache/test_dllm_fdfo_kv_reuse.py index fb9879a13..6152e59cc 100644 --- a/test/registered/unit/mem_cache/test_dllm_fdfo_kv_reuse.py +++ b/test/registered/unit/mem_cache/test_dllm_fdfo_kv_reuse.py @@ -7,6 +7,7 @@ from types import SimpleNamespace import torch from sglang.srt.dllm.mixin.scheduler import DllmManager +from sglang.srt.managers.schedule_batch import ReqKvInfo from sglang.srt.mem_cache.allocation import alloc_for_extend from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.runtime_context import get_context @@ -63,7 +64,7 @@ def _make_req(rid, prefix, block_size, *, req_pool_idx=None, reuse=False): prefix_indices=torch.tensor(prefix, dtype=torch.int32), dllm_incomplete_ids=array("q", range(block_size)) if reuse else array("q"), inflight_middle_chunks=1 if req_pool_idx is not None else 0, - kv=SimpleNamespace( + kv=ReqKvInfo( req_pool_idx=req_pool_idx, kv_committed_len=len(prefix) if req_pool_idx is not None else 0, kv_allocated_len=(