[logprob] Fix logprob + streaming for long concurrent decode by caching already processed logprob (#17005)
Co-authored-by: root <root@memx-cge-29-sr1.xpop.twttr.net>
This commit is contained in:
@@ -165,6 +165,14 @@ class ReqState:
|
||||
output_token_ids_logprobs_val: List = dataclasses.field(default_factory=list)
|
||||
output_token_ids_logprobs_idx: List = dataclasses.field(default_factory=list)
|
||||
|
||||
# For detokenized logprobs
|
||||
input_token_logprobs: List[Any] = dataclasses.field(default_factory=list)
|
||||
output_token_logprobs: List[Any] = dataclasses.field(default_factory=list)
|
||||
input_top_logprobs: List[Any] = dataclasses.field(default_factory=list)
|
||||
output_top_logprobs: List[Any] = dataclasses.field(default_factory=list)
|
||||
input_token_ids_logprobs: List[Any] = dataclasses.field(default_factory=list)
|
||||
output_token_ids_logprobs: List[Any] = dataclasses.field(default_factory=list)
|
||||
|
||||
|
||||
class InputFormat(Enum):
|
||||
"""Input format types for tokenization handling."""
|
||||
@@ -1616,43 +1624,84 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
||||
token_ids_logprob: List[int],
|
||||
return_text_in_logprobs: bool,
|
||||
):
|
||||
meta_info["input_token_logprobs"] = self.detokenize_logprob_tokens(
|
||||
state.input_token_logprobs_val,
|
||||
state.input_token_logprobs_idx,
|
||||
return_text_in_logprobs,
|
||||
)
|
||||
meta_info["output_token_logprobs"] = self.detokenize_logprob_tokens(
|
||||
state.output_token_logprobs_val,
|
||||
state.output_token_logprobs_idx,
|
||||
return_text_in_logprobs,
|
||||
)
|
||||
|
||||
if top_logprobs_num > 0:
|
||||
meta_info["input_top_logprobs"] = self.detokenize_top_logprobs_tokens(
|
||||
state.input_top_logprobs_val,
|
||||
state.input_top_logprobs_idx,
|
||||
return_text_in_logprobs,
|
||||
)
|
||||
meta_info["output_top_logprobs"] = self.detokenize_top_logprobs_tokens(
|
||||
state.output_top_logprobs_val,
|
||||
state.output_top_logprobs_idx,
|
||||
return_text_in_logprobs,
|
||||
)
|
||||
|
||||
if token_ids_logprob is not None:
|
||||
meta_info["input_token_ids_logprobs"] = self.detokenize_top_logprobs_tokens(
|
||||
state.input_token_ids_logprobs_val,
|
||||
state.input_token_ids_logprobs_idx,
|
||||
return_text_in_logprobs,
|
||||
)
|
||||
meta_info["output_token_ids_logprobs"] = (
|
||||
self.detokenize_top_logprobs_tokens(
|
||||
state.output_token_ids_logprobs_val,
|
||||
state.output_token_ids_logprobs_idx,
|
||||
# 1. Handle regular logprobs
|
||||
if len(state.input_token_logprobs_val) > len(state.input_token_logprobs):
|
||||
state.input_token_logprobs.extend(
|
||||
self.detokenize_logprob_tokens(
|
||||
state.input_token_logprobs_val[len(state.input_token_logprobs) :],
|
||||
state.input_token_logprobs_idx[len(state.input_token_logprobs) :],
|
||||
return_text_in_logprobs,
|
||||
)
|
||||
)
|
||||
|
||||
if len(state.output_token_logprobs_val) > len(state.output_token_logprobs):
|
||||
state.output_token_logprobs.extend(
|
||||
self.detokenize_logprob_tokens(
|
||||
state.output_token_logprobs_val[len(state.output_token_logprobs) :],
|
||||
state.output_token_logprobs_idx[len(state.output_token_logprobs) :],
|
||||
return_text_in_logprobs,
|
||||
)
|
||||
)
|
||||
|
||||
meta_info["input_token_logprobs"] = state.input_token_logprobs
|
||||
meta_info["output_token_logprobs"] = state.output_token_logprobs
|
||||
|
||||
# 2. Handle top logprobs
|
||||
if top_logprobs_num > 0:
|
||||
if len(state.input_top_logprobs_val) > len(state.input_top_logprobs):
|
||||
state.input_top_logprobs.extend(
|
||||
self.detokenize_top_logprobs_tokens(
|
||||
state.input_top_logprobs_val[len(state.input_top_logprobs) :],
|
||||
state.input_top_logprobs_idx[len(state.input_top_logprobs) :],
|
||||
return_text_in_logprobs,
|
||||
)
|
||||
)
|
||||
if len(state.output_top_logprobs_val) > len(state.output_top_logprobs):
|
||||
state.output_top_logprobs.extend(
|
||||
self.detokenize_top_logprobs_tokens(
|
||||
state.output_top_logprobs_val[len(state.output_top_logprobs) :],
|
||||
state.output_top_logprobs_idx[len(state.output_top_logprobs) :],
|
||||
return_text_in_logprobs,
|
||||
)
|
||||
)
|
||||
|
||||
meta_info["input_top_logprobs"] = state.input_top_logprobs
|
||||
meta_info["output_top_logprobs"] = state.output_top_logprobs
|
||||
|
||||
# 3. Handle token_ids_logprob
|
||||
if token_ids_logprob is not None:
|
||||
if len(state.input_token_ids_logprobs_val) > len(
|
||||
state.input_token_ids_logprobs
|
||||
):
|
||||
state.input_token_ids_logprobs.extend(
|
||||
self.detokenize_top_logprobs_tokens(
|
||||
state.input_token_ids_logprobs_val[
|
||||
len(state.input_token_ids_logprobs) :
|
||||
],
|
||||
state.input_token_ids_logprobs_idx[
|
||||
len(state.input_token_ids_logprobs) :
|
||||
],
|
||||
return_text_in_logprobs,
|
||||
)
|
||||
)
|
||||
if len(state.output_token_ids_logprobs_val) > len(
|
||||
state.output_token_ids_logprobs
|
||||
):
|
||||
state.output_token_ids_logprobs.extend(
|
||||
self.detokenize_top_logprobs_tokens(
|
||||
state.output_token_ids_logprobs_val[
|
||||
len(state.output_token_ids_logprobs) :
|
||||
],
|
||||
state.output_token_ids_logprobs_idx[
|
||||
len(state.output_token_ids_logprobs) :
|
||||
],
|
||||
return_text_in_logprobs,
|
||||
)
|
||||
)
|
||||
|
||||
meta_info["input_token_ids_logprobs"] = state.input_token_ids_logprobs
|
||||
meta_info["output_token_ids_logprobs"] = state.output_token_ids_logprobs
|
||||
|
||||
def convert_logprob_style(
|
||||
self,
|
||||
meta_info: dict,
|
||||
|
||||
Reference in New Issue
Block a user