Unify memory management across (overlap, non-overlap) x (page>=1) x (spec, non-spec, spec v2) x (retract, finished) (#12224)

This commit is contained in:
Liangsheng Yin
2025-11-11 02:56:22 +08:00
committed by GitHub
parent 838bcb0d93
commit 665416f6dd
24 changed files with 193 additions and 156 deletions
@@ -116,6 +116,8 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
len(batch.input_ids),
)
end_offset = batch.seq_lens + self.draft_token_num
for req in batch.reqs:
req.kv_allocated_len += 1
else:
prefix_lens = batch.seq_lens
prefix_lens_cpu = batch.seq_lens_cpu
@@ -415,6 +417,9 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
if page_size == 1:
# TODO: boolean array index leads to a device sync. Remove it.
token_to_kv_pool_allocator.free(batch.out_cache_loc[evict_mask])
for i, req in enumerate(batch.reqs):
req.kv_committed_len += accept_length_list[i] + 1
req.kv_allocated_len = req.kv_committed_len
else:
if self.topk == 1:
# Only evict full empty page. Do not evict partial empty page
@@ -426,6 +431,9 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
next_power_of_2(self.draft_token_num),
)
token_to_kv_pool_allocator.free(batch.out_cache_loc[evict_mask])
for i, req in enumerate(batch.reqs):
req.kv_committed_len += accept_length_list[i] + 1
req.kv_allocated_len = req.kv_committed_len
else:
# Shift the accepted tokens to the beginning.
# Only evict the last part
@@ -129,6 +129,10 @@ class EagleDraftInputV2Mixin:
batch.seq_lens_cpu = batch.seq_lens.cpu()
batch.seq_lens_sum = batch.seq_lens_cpu.sum().item()
for i, req in enumerate(batch.reqs):
req.kv_committed_len = batch.seq_lens_cpu[i].item()
req.kv_allocated_len = req.kv_committed_len + self.ALLOC_LEN_PER_DECODE
def prepare_for_v2_draft(
self: EagleDraftInput,
req_to_token_pool: ReqToTokenPool,
@@ -364,6 +364,8 @@ class EAGLEWorker(TpModelWorker):
# [ topk 0 ] [ topk 1 ]
# [iter=0, iter=1, iter=2] [iter=0, iter=1, iter=2]
if self.page_size == 1:
for req in batch.reqs:
req.kv_allocated_len += self.speculative_num_steps * self.topk
out_cache_loc, token_to_kv_pool_state_backup = alloc_token_slots(
batch.tree_cache,
num_seqs * self.speculative_num_steps * self.topk,
+10 -2
View File
@@ -195,7 +195,9 @@ class NgramVerifyInput(SpecInput):
logits_output.hidden_states = logits_output.hidden_states[self.accept_index]
self.verified_id = self.predict[self.accept_index]
def _free_cache(self, batch: ScheduleBatch, page_size: int):
def _free_cache(
self, batch: ScheduleBatch, page_size: int, accept_length_cpu: torch.Tensor
):
bs = batch.batch_size()
# Free the KV cache for unaccepted tokens
if page_size == 1:
@@ -250,6 +252,11 @@ class NgramVerifyInput(SpecInput):
)
batch.out_cache_loc = tgt_cache_loc
accept_length_list = accept_length_cpu.tolist()
for i, req in enumerate(batch.reqs):
req.kv_committed_len += accept_length_list[i] + 1
req.kv_allocated_len = req.kv_committed_len
assign_req_to_token_pool[(bs,)](
batch.req_pool_indices,
batch.req_to_token_pool.req_to_token,
@@ -416,11 +423,12 @@ class NgramVerifyInput(SpecInput):
# self._sampling_verify(batch, logits_output, sampling_info)
self._fill_requests(batch, logits_output)
self._free_cache(batch, page_size)
accept_length_cpu = self.accept_length.cpu()
num_accepted_tokens = accept_length_cpu.sum().item()
self._free_cache(batch, page_size, accept_length_cpu)
batch.seq_lens.add_(self.accept_length + 1)
batch.seq_lens_cpu.add_(accept_length_cpu + 1)