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:
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user