Pull the max-prefix-len computation into its own helper and rename the matched-token argument (#25728)
This commit is contained in:
@@ -1026,25 +1026,21 @@ class Req(ReqDllmMixin):
|
|||||||
)
|
)
|
||||||
self.logprob_start_len = -1
|
self.logprob_start_len = -1
|
||||||
|
|
||||||
# NOTE: the matched length is at most 1 less than the input length to enable logprob computation
|
token_ids_to_match = self.fill_ids[: self._compute_max_prefix_len(input_len)]
|
||||||
max_prefix_len = input_len - 1
|
|
||||||
if self.return_logprob and self.logprob_start_len >= 0:
|
|
||||||
max_prefix_len = min(max_prefix_len, self.logprob_start_len)
|
|
||||||
max_prefix_len = max(max_prefix_len, 0)
|
|
||||||
token_ids = self.fill_ids[:max_prefix_len]
|
|
||||||
del max_prefix_len
|
|
||||||
|
|
||||||
# Disable prefix caching when embed overrides are present: same token IDs
|
# Disable prefix caching when embed overrides are present: same token IDs
|
||||||
# with different override vectors must not share cached KV values.
|
# with different override vectors must not share cached KV values.
|
||||||
if self.positional_embed_overrides is not None:
|
if self.positional_embed_overrides is not None:
|
||||||
token_ids = []
|
token_ids_to_match = []
|
||||||
|
|
||||||
if tree_cache is not None:
|
if tree_cache is not None:
|
||||||
if cow_mamba is None:
|
if cow_mamba is None:
|
||||||
cow_mamba = tree_cache.supports_mamba()
|
cow_mamba = tree_cache.supports_mamba()
|
||||||
match_result = tree_cache.match_prefix(
|
match_result = tree_cache.match_prefix(
|
||||||
MatchPrefixParams(
|
MatchPrefixParams(
|
||||||
key=RadixKey(token_ids=token_ids, extra_key=self.extra_key),
|
key=RadixKey(
|
||||||
|
token_ids=token_ids_to_match, extra_key=self.extra_key
|
||||||
|
),
|
||||||
req=self,
|
req=self,
|
||||||
cow_mamba=cow_mamba,
|
cow_mamba=cow_mamba,
|
||||||
)
|
)
|
||||||
@@ -1091,6 +1087,13 @@ class Req(ReqDllmMixin):
|
|||||||
|
|
||||||
self.set_extend_input_len(len(self.fill_ids) - len(self.prefix_indices))
|
self.set_extend_input_len(len(self.fill_ids) - 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
|
||||||
|
if self.return_logprob and self.logprob_start_len >= 0:
|
||||||
|
max_prefix_len = min(max_prefix_len, self.logprob_start_len)
|
||||||
|
return max(max_prefix_len, 0)
|
||||||
|
|
||||||
# Based on https://github.com/vllm-project/vllm/blob/7a64d24aad69e4d2548aa0bf528d9fe63428ab01/vllm/transformers_utils/detokenizer.py#L194-L313
|
# Based on https://github.com/vllm-project/vllm/blob/7a64d24aad69e4d2548aa0bf528d9fe63428ab01/vllm/transformers_utils/detokenizer.py#L194-L313
|
||||||
def init_incremental_detokenize(self):
|
def init_incremental_detokenize(self):
|
||||||
first_iter = self.surr_offset is None or self.read_offset is None
|
first_iter = self.surr_offset is None or self.read_offset is None
|
||||||
|
|||||||
Reference in New Issue
Block a user