Avoid dual semantics of extend_input_len by computing the candidate on the fly (#27616)
This commit is contained in:
@@ -1218,8 +1218,6 @@ class Req(ReqDllmMixin):
|
||||
)
|
||||
)
|
||||
|
||||
self.set_extend_input_len(input_len - len(self.prefix_indices))
|
||||
|
||||
def _compute_max_prefix_len(self, input_len: int) -> int:
|
||||
# NOTE: the matched length is at most 1 less than the input length to enable logprob computation
|
||||
max_prefix_len = input_len - 1
|
||||
|
||||
@@ -659,7 +659,7 @@ class PrefillAdder:
|
||||
* self.page_size
|
||||
)
|
||||
|
||||
req.extend_input_len = trunc_len
|
||||
req.set_extend_input_len(trunc_len)
|
||||
req.fill_len = prefix_len + trunc_len
|
||||
|
||||
self.can_run_list.append(req)
|
||||
@@ -679,9 +679,13 @@ class PrefillAdder:
|
||||
return AddReqResult.NO_TOKEN
|
||||
|
||||
# Truncate input length to available tokens and update request metadata
|
||||
truncated = req.extend_input_len > _rem_tokens
|
||||
req.extend_input_len = min(req.extend_input_len, _rem_tokens)
|
||||
req.fill_len = len(req.prefix_indices) + req.extend_input_len
|
||||
cand_extend_input_len = len(req.full_untruncated_fill_ids) - len(
|
||||
req.prefix_indices
|
||||
)
|
||||
truncated = cand_extend_input_len > _rem_tokens
|
||||
new_len = min(cand_extend_input_len, _rem_tokens)
|
||||
req.set_extend_input_len(new_len)
|
||||
req.fill_len = len(req.prefix_indices) + new_len
|
||||
self.can_run_list.append(req)
|
||||
|
||||
# Update budget: reserve max_new_tokens only if not truncated
|
||||
@@ -719,9 +723,13 @@ class PrefillAdder:
|
||||
return req
|
||||
_rem_tokens = self.rem_chunk_tokens
|
||||
|
||||
truncated = req.extend_input_len > _rem_tokens
|
||||
req.set_extend_input_len(min(req.extend_input_len, _rem_tokens))
|
||||
req.fill_len = len(req.prefix_indices) + req.extend_input_len
|
||||
cand_extend_input_len = len(req.full_untruncated_fill_ids) - len(
|
||||
req.prefix_indices
|
||||
)
|
||||
truncated = cand_extend_input_len > _rem_tokens
|
||||
new_len = min(cand_extend_input_len, _rem_tokens)
|
||||
req.set_extend_input_len(new_len)
|
||||
req.fill_len = len(req.prefix_indices) + new_len
|
||||
self.can_run_list.append(req)
|
||||
self._update_prefill_budget(
|
||||
0,
|
||||
@@ -755,11 +763,14 @@ class PrefillAdder:
|
||||
self.tree_cache.dec_lock_ref(last_node)
|
||||
|
||||
def add_one_req_ignore_eos(self, req: Req):
|
||||
paged_input = self.ceil_paged_tokens(req.extend_input_len)
|
||||
cand_extend_input_len = len(req.full_untruncated_fill_ids) - len(
|
||||
req.prefix_indices
|
||||
)
|
||||
paged_input = self.ceil_paged_tokens(cand_extend_input_len)
|
||||
if paged_input > min(self.cur_rem_tokens, self.rem_total_tokens):
|
||||
return AddReqResult.NO_TOKEN
|
||||
if self.is_hybrid_swa:
|
||||
if self._swa_budget_for_req(req.extend_input_len) > self.rem_swa_tokens:
|
||||
if self._swa_budget_for_req(cand_extend_input_len) > self.rem_swa_tokens:
|
||||
return AddReqResult.NO_TOKEN
|
||||
|
||||
def add_req_state(r, insert_sort=False):
|
||||
@@ -799,7 +810,7 @@ class PrefillAdder:
|
||||
# Skip this logic for swa. The SWA has different memory management, and
|
||||
# this mechanism is underestimating the memory usage.
|
||||
cur_rem_tokens = self.cur_rem_tokens - self.ceil_paged_tokens(
|
||||
req.extend_input_len
|
||||
cand_extend_input_len
|
||||
)
|
||||
tokens_freed = 0
|
||||
for i, (tokens_left, tokens_occupied) in enumerate(self.req_states):
|
||||
@@ -825,13 +836,13 @@ class PrefillAdder:
|
||||
self._add_dllm_req(req, 0)
|
||||
elif (
|
||||
self.rem_chunk_tokens is None # chunked prefill is disabled
|
||||
or req.extend_input_len <= self.rem_chunk_tokens # it is the last chunk
|
||||
or cand_extend_input_len <= self.rem_chunk_tokens # it is the last chunk
|
||||
):
|
||||
# Non-chunked prefill — the whole sequence is committed this iter.
|
||||
req.set_extend_input_len(
|
||||
len(req.full_untruncated_fill_ids) - len(req.prefix_indices)
|
||||
)
|
||||
req.fill_len = len(req.full_untruncated_fill_ids)
|
||||
assert (
|
||||
req.fill_len == len(req.prefix_indices) + req.extend_input_len
|
||||
), f"{req.fill_len=} {len(req.prefix_indices)=} {req.extend_input_len=}"
|
||||
self.can_run_list.append(req)
|
||||
self._update_prefill_budget(
|
||||
0,
|
||||
@@ -889,10 +900,13 @@ class PrefillAdder:
|
||||
max(req.sampling_params.max_new_tokens - len(req.output_ids), 0),
|
||||
CLIP_MAX_NEW_TOKENS,
|
||||
)
|
||||
total_tokens = req.extend_input_len + max_new + self.page_size
|
||||
cand_extend_input_len = len(req.full_untruncated_fill_ids) - len(
|
||||
req.prefix_indices
|
||||
)
|
||||
total_tokens = cand_extend_input_len + max_new + self.page_size
|
||||
|
||||
# adjusting the input_tokens based on host_hit_length and page_size
|
||||
real_input_tokens = req.extend_input_len - req.host_hit_length
|
||||
real_input_tokens = cand_extend_input_len - req.host_hit_length
|
||||
real_input_tokens = self.ceil_paged_tokens(real_input_tokens)
|
||||
prefix_len = len(req.prefix_indices)
|
||||
|
||||
@@ -901,7 +915,7 @@ class PrefillAdder:
|
||||
|
||||
if self.is_hybrid_swa:
|
||||
swa_needed = self._swa_budget_for_req(
|
||||
req.extend_input_len, swa_host_hit_length=req.swa_host_hit_length
|
||||
cand_extend_input_len, swa_host_hit_length=req.swa_host_hit_length
|
||||
)
|
||||
if swa_needed >= self.rem_swa_tokens:
|
||||
return AddReqResult.NO_TOKEN
|
||||
@@ -923,7 +937,7 @@ class PrefillAdder:
|
||||
|
||||
if self.is_hybrid_swa:
|
||||
swa_needed = self._swa_budget_for_req(
|
||||
req.extend_input_len, swa_host_hit_length=req.swa_host_hit_length
|
||||
cand_extend_input_len, swa_host_hit_length=req.swa_host_hit_length
|
||||
)
|
||||
if swa_needed >= self.rem_swa_tokens:
|
||||
return AddReqResult.NO_TOKEN
|
||||
@@ -937,13 +951,12 @@ class PrefillAdder:
|
||||
)
|
||||
)
|
||||
req.prefix_indices = torch.cat([req.prefix_indices, new_indices])
|
||||
req.set_extend_input_len(
|
||||
len(req.full_untruncated_fill_ids) - len(req.prefix_indices)
|
||||
)
|
||||
prefix_len = len(req.prefix_indices)
|
||||
req.cache_protected_len = prefix_len
|
||||
|
||||
input_tokens = self.ceil_paged_tokens(req.extend_input_len)
|
||||
input_tokens = self.ceil_paged_tokens(
|
||||
len(req.full_untruncated_fill_ids) - len(req.prefix_indices)
|
||||
)
|
||||
|
||||
if (
|
||||
self.rem_chunk_tokens is None
|
||||
@@ -967,10 +980,10 @@ class PrefillAdder:
|
||||
self._req_inc_lock_ref(req)
|
||||
elif self.rem_chunk_tokens is None or input_tokens <= self.rem_chunk_tokens:
|
||||
# Non-chunked prefill — the whole sequence is committed this iter.
|
||||
req.set_extend_input_len(
|
||||
len(req.full_untruncated_fill_ids) - len(req.prefix_indices)
|
||||
)
|
||||
req.fill_len = len(req.full_untruncated_fill_ids)
|
||||
assert (
|
||||
req.fill_len == len(req.prefix_indices) + req.extend_input_len
|
||||
), f"{req.fill_len=} {len(req.prefix_indices)=} {req.extend_input_len=}"
|
||||
self.can_run_list.append(req)
|
||||
|
||||
self._req_inc_lock_ref(req)
|
||||
@@ -1052,7 +1065,8 @@ class PrefillAdder:
|
||||
|
||||
preemptible_reqs = []
|
||||
min_tokens_to_remove = (
|
||||
req.extend_input_len
|
||||
len(req.full_untruncated_fill_ids)
|
||||
- len(req.prefix_indices)
|
||||
+ min(req.sampling_params.max_new_tokens, CLIP_MAX_NEW_TOKENS)
|
||||
- self.rem_total_tokens
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user