[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:
Byron Hsu
2026-01-13 12:39:14 -08:00
committed by GitHub
parent 075c5a5789
commit 339915ce2b

View File

@@ -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,