Reserve slot 0 as padding in all req pools (#24243)
This commit is contained in:
@@ -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()
|
||||
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user