diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index b33315b8f..8644e8192 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -130,13 +130,14 @@ class DecodeReqToTokenPool: self.device = device self.pre_alloc_size = pre_alloc_size with memory_saver_adapter.region(tag=GPU_MEMORY_TYPE_KV_CACHE): + # +1 row 0 padding; mirrors ReqToTokenPool / KV pool padding slot 0. self.req_to_token = torch.zeros( - (size + pre_alloc_size, max_context_len), + (size + pre_alloc_size + 1, max_context_len), dtype=torch.int32, device=device, ) - self.free_slots = list(range(size + pre_alloc_size)) + self.free_slots = list(range(1, size + pre_alloc_size + 1)) def write(self, indices, values): self.req_to_token[indices] = values @@ -173,7 +174,7 @@ class DecodeReqToTokenPool: req.req_pool_idx = None def clear(self): - self.free_slots = list(range(self.size + self.pre_alloc_size)) + self.free_slots = list(range(1, self.size + self.pre_alloc_size + 1)) class HybridMambaDecodeReqToTokenPool(HybridReqToTokenPool): @@ -238,7 +239,7 @@ class HybridMambaDecodeReqToTokenPool(HybridReqToTokenPool): ) def clear(self): - self.free_slots = list(range(self.size + self.pre_alloc_size)) + self.free_slots = list(range(1, self.size + self.pre_alloc_size + 1)) self.mamba_pool.clear() diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index bbc25a7e6..a0e0765f2 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -142,10 +142,14 @@ class ReqToTokenPool: self.max_context_len = max_context_len self.device = device with memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE): + # +1 row for padding slot 0 (mirrors KV pool): cuda-graph padded + # batches default req_pool_indices to 0, so routing dummies through + # unowned slot 0 keeps req_to_token[0, :] zero and downstream writes + # harmless. self.req_to_token = torch.zeros( - (size, max_context_len), dtype=torch.int32, device=device + (size + 1, max_context_len), dtype=torch.int32, device=device ) - self.free_slots = list(range(size)) + self.free_slots = list(range(1, size + 1)) def write(self, indices, values): self.req_to_token[indices] = values @@ -185,7 +189,7 @@ class ReqToTokenPool: req.req_pool_idx = None def clear(self): - self.free_slots = list(range(self.size)) + self.free_slots = list(range(1, self.size + 1)) class MambaPool: