Avoid dual semantics of extend_input_len by computing the candidate on the fly (#27616)

This commit is contained in:
fzyzcjy
2026-06-25 08:17:37 +08:00
committed by GitHub
parent 5d4e63d49e
commit 563c3418a7
3 changed files with 42 additions and 28 deletions
@@ -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
+40 -26
View File
@@ -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)