From 339915ce2b8c9c6cb73a4a52d3d084d8f66c967f Mon Sep 17 00:00:00 2001 From: Byron Hsu Date: Tue, 13 Jan 2026 12:39:14 -0800 Subject: [PATCH] [logprob] Fix logprob + streaming for long concurrent decode by caching already processed logprob (#17005) Co-authored-by: root --- .../sglang/srt/managers/tokenizer_manager.py | 115 +++++++++++++----- 1 file changed, 82 insertions(+), 33 deletions(-) diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index a433a0597..a3c5001e8 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -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,