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:
|
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
|
# NOTE: the matched length is at most 1 less than the input length to enable logprob computation
|
||||||
max_prefix_len = input_len - 1
|
max_prefix_len = input_len - 1
|
||||||
|
|||||||
@@ -659,7 +659,7 @@ class PrefillAdder:
|
|||||||
* self.page_size
|
* self.page_size
|
||||||
)
|
)
|
||||||
|
|
||||||
req.extend_input_len = trunc_len
|
req.set_extend_input_len(trunc_len)
|
||||||
req.fill_len = prefix_len + trunc_len
|
req.fill_len = prefix_len + trunc_len
|
||||||
|
|
||||||
self.can_run_list.append(req)
|
self.can_run_list.append(req)
|
||||||
@@ -679,9 +679,13 @@ class PrefillAdder:
|
|||||||
return AddReqResult.NO_TOKEN
|
return AddReqResult.NO_TOKEN
|
||||||
|
|
||||||
# Truncate input length to available tokens and update request metadata
|
# Truncate input length to available tokens and update request metadata
|
||||||
truncated = req.extend_input_len > _rem_tokens
|
cand_extend_input_len = len(req.full_untruncated_fill_ids) - len(
|
||||||
req.extend_input_len = min(req.extend_input_len, _rem_tokens)
|
req.prefix_indices
|
||||||
req.fill_len = len(req.prefix_indices) + req.extend_input_len
|
)
|
||||||
|
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.can_run_list.append(req)
|
||||||
|
|
||||||
# Update budget: reserve max_new_tokens only if not truncated
|
# Update budget: reserve max_new_tokens only if not truncated
|
||||||
@@ -719,9 +723,13 @@ class PrefillAdder:
|
|||||||
return req
|
return req
|
||||||
_rem_tokens = self.rem_chunk_tokens
|
_rem_tokens = self.rem_chunk_tokens
|
||||||
|
|
||||||
truncated = req.extend_input_len > _rem_tokens
|
cand_extend_input_len = len(req.full_untruncated_fill_ids) - len(
|
||||||
req.set_extend_input_len(min(req.extend_input_len, _rem_tokens))
|
req.prefix_indices
|
||||||
req.fill_len = len(req.prefix_indices) + req.extend_input_len
|
)
|
||||||
|
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.can_run_list.append(req)
|
||||||
self._update_prefill_budget(
|
self._update_prefill_budget(
|
||||||
0,
|
0,
|
||||||
@@ -755,11 +763,14 @@ class PrefillAdder:
|
|||||||
self.tree_cache.dec_lock_ref(last_node)
|
self.tree_cache.dec_lock_ref(last_node)
|
||||||
|
|
||||||
def add_one_req_ignore_eos(self, req: Req):
|
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):
|
if paged_input > min(self.cur_rem_tokens, self.rem_total_tokens):
|
||||||
return AddReqResult.NO_TOKEN
|
return AddReqResult.NO_TOKEN
|
||||||
if self.is_hybrid_swa:
|
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
|
return AddReqResult.NO_TOKEN
|
||||||
|
|
||||||
def add_req_state(r, insert_sort=False):
|
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
|
# Skip this logic for swa. The SWA has different memory management, and
|
||||||
# this mechanism is underestimating the memory usage.
|
# this mechanism is underestimating the memory usage.
|
||||||
cur_rem_tokens = self.cur_rem_tokens - self.ceil_paged_tokens(
|
cur_rem_tokens = self.cur_rem_tokens - self.ceil_paged_tokens(
|
||||||
req.extend_input_len
|
cand_extend_input_len
|
||||||
)
|
)
|
||||||
tokens_freed = 0
|
tokens_freed = 0
|
||||||
for i, (tokens_left, tokens_occupied) in enumerate(self.req_states):
|
for i, (tokens_left, tokens_occupied) in enumerate(self.req_states):
|
||||||
@@ -825,13 +836,13 @@ class PrefillAdder:
|
|||||||
self._add_dllm_req(req, 0)
|
self._add_dllm_req(req, 0)
|
||||||
elif (
|
elif (
|
||||||
self.rem_chunk_tokens is None # chunked prefill is disabled
|
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.
|
# 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)
|
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.can_run_list.append(req)
|
||||||
self._update_prefill_budget(
|
self._update_prefill_budget(
|
||||||
0,
|
0,
|
||||||
@@ -889,10 +900,13 @@ class PrefillAdder:
|
|||||||
max(req.sampling_params.max_new_tokens - len(req.output_ids), 0),
|
max(req.sampling_params.max_new_tokens - len(req.output_ids), 0),
|
||||||
CLIP_MAX_NEW_TOKENS,
|
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
|
# 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)
|
real_input_tokens = self.ceil_paged_tokens(real_input_tokens)
|
||||||
prefix_len = len(req.prefix_indices)
|
prefix_len = len(req.prefix_indices)
|
||||||
|
|
||||||
@@ -901,7 +915,7 @@ class PrefillAdder:
|
|||||||
|
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
swa_needed = self._swa_budget_for_req(
|
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:
|
if swa_needed >= self.rem_swa_tokens:
|
||||||
return AddReqResult.NO_TOKEN
|
return AddReqResult.NO_TOKEN
|
||||||
@@ -923,7 +937,7 @@ class PrefillAdder:
|
|||||||
|
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
swa_needed = self._swa_budget_for_req(
|
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:
|
if swa_needed >= self.rem_swa_tokens:
|
||||||
return AddReqResult.NO_TOKEN
|
return AddReqResult.NO_TOKEN
|
||||||
@@ -937,13 +951,12 @@ class PrefillAdder:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
req.prefix_indices = torch.cat([req.prefix_indices, new_indices])
|
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)
|
prefix_len = len(req.prefix_indices)
|
||||||
req.cache_protected_len = prefix_len
|
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 (
|
if (
|
||||||
self.rem_chunk_tokens is None
|
self.rem_chunk_tokens is None
|
||||||
@@ -967,10 +980,10 @@ class PrefillAdder:
|
|||||||
self._req_inc_lock_ref(req)
|
self._req_inc_lock_ref(req)
|
||||||
elif self.rem_chunk_tokens is None or input_tokens <= self.rem_chunk_tokens:
|
elif self.rem_chunk_tokens is None or input_tokens <= self.rem_chunk_tokens:
|
||||||
# Non-chunked prefill — the whole sequence is committed this iter.
|
# 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)
|
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.can_run_list.append(req)
|
||||||
|
|
||||||
self._req_inc_lock_ref(req)
|
self._req_inc_lock_ref(req)
|
||||||
@@ -1052,7 +1065,8 @@ class PrefillAdder:
|
|||||||
|
|
||||||
preemptible_reqs = []
|
preemptible_reqs = []
|
||||||
min_tokens_to_remove = (
|
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)
|
+ min(req.sampling_params.max_new_tokens, CLIP_MAX_NEW_TOKENS)
|
||||||
- self.rem_total_tokens
|
- self.rem_total_tokens
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -77,6 +77,8 @@ class TestPrefillAdder(CustomTestCase):
|
|||||||
req.rid = str(rid)
|
req.rid = str(rid)
|
||||||
req.priority = priority
|
req.priority = priority
|
||||||
req.extend_input_len = 0
|
req.extend_input_len = 0
|
||||||
|
req.prefix_indices = []
|
||||||
|
req.full_untruncated_fill_ids = []
|
||||||
req.extend_logprob_start_len = 0
|
req.extend_logprob_start_len = 0
|
||||||
req.output_ids = [0] * output_len
|
req.output_ids = [0] * output_len
|
||||||
req.sampling_params = SimpleNamespace(max_new_tokens=max_new_tokens)
|
req.sampling_params = SimpleNamespace(max_new_tokens=max_new_tokens)
|
||||||
|
|||||||
Reference in New Issue
Block a user