From e737f61b297e4d16ccb6196306349fe4bcbf4b35 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Mon, 18 May 2026 18:42:59 +0800 Subject: [PATCH] Introduce SchedulerLogprobResultProcessor to own logprob state (#25632) --- python/sglang/srt/disaggregation/prefill.py | 2 + python/sglang/srt/managers/scheduler.py | 8 ++ .../logprob_result_processor.py | 13 +++ .../scheduler_output_processor_mixin.py | 103 +++++++++++++----- 4 files changed, 100 insertions(+), 26 deletions(-) create mode 100644 python/sglang/srt/managers/scheduler_components/logprob_result_processor.py diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 732e78453..c7d191702 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -536,6 +536,7 @@ class SchedulerDisaggregationPrefillMixin: extend_input_len = extend_input_len_per_req[i] num_input_logprobs = extend_input_len - extend_logprob_start_len self.add_logprob_return_values( + self.logprob_result_processor, i, req, logprob_pt, @@ -573,6 +574,7 @@ class SchedulerDisaggregationPrefillMixin: # Update input logprobs. num_input_logprobs = extend_input_len - extend_logprob_start_len self.add_input_logprob_return_values( + self.logprob_result_processor, i, req, logits_output, diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 0e75c36eb..531e76dd1 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -176,6 +176,9 @@ from sglang.srt.managers.scheduler_components.kv_events_publisher import ( from sglang.srt.managers.scheduler_components.load_inquirer import ( SchedulerLoadInquirer, ) +from sglang.srt.managers.scheduler_components.logprob_result_processor import ( + SchedulerLogprobResultProcessor, +) from sglang.srt.managers.scheduler_components.metrics_reporter import ( RECORD_STEP_TIME, PrefillStats, @@ -734,6 +737,11 @@ class Scheduler( get_spec_total_num_forward_ct=lambda: self.metrics_reporter.spec_total_num_forward_ct, ) + self.logprob_result_processor = SchedulerLogprobResultProcessor( + server_args=self.server_args, + model_config=self.model_config, + ) + self.is_initializing = False def init_zbal_on_npu(self): diff --git a/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py b/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py new file mode 100644 index 000000000..c8a5427f3 --- /dev/null +++ b/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py @@ -0,0 +1,13 @@ +from __future__ import annotations + +from dataclasses import dataclass + + +from sglang.srt.configs.model_config import ModelConfig +from sglang.srt.server_args import ServerArgs + + +@dataclass(kw_only=True, slots=True, frozen=True) +class SchedulerLogprobResultProcessor: + server_args: ServerArgs + model_config: ModelConfig diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index 34551b1d2..06f313df8 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -33,6 +33,9 @@ if TYPE_CHECKING: ScheduleBatch, Scheduler, ) + from sglang.srt.managers.scheduler_components.logprob_result_processor import ( + SchedulerLogprobResultProcessor, + ) logger = logging.getLogger(__name__) @@ -268,12 +271,16 @@ class SchedulerOutputProcessorMixin: extend_logprob_start_len = extend_logprob_start_len_per_req[i] extend_input_len = extend_input_len_per_req[i] - num_input_logprobs = self._calculate_num_input_logprobs( - req, extend_input_len, extend_logprob_start_len + num_input_logprobs = self.calculate_num_input_logprobs( + self.logprob_result_processor, + req, + extend_input_len, + extend_logprob_start_len, ) if req.return_logprob: self.add_logprob_return_values( + self.logprob_result_processor, i, req, logprob_pt, @@ -326,11 +333,15 @@ class SchedulerOutputProcessorMixin: extend_input_len = extend_input_len_per_req[i] if extend_logprob_start_len < extend_input_len: # Update input logprobs. - num_input_logprobs = self._calculate_num_input_logprobs( - req, extend_input_len, extend_logprob_start_len + num_input_logprobs = self.calculate_num_input_logprobs( + self.logprob_result_processor, + req, + extend_input_len, + extend_logprob_start_len, ) if req.return_logprob: self.add_input_logprob_return_values( + self.logprob_result_processor, i, req, logits_output, @@ -709,11 +720,14 @@ class SchedulerOutputProcessorMixin: actual_seq_len // mamba_track_interval * mamba_track_interval ) + @staticmethod def _process_input_token_logprobs( - self: Scheduler, req: Req, input_token_logprobs: List + self: "SchedulerLogprobResultProcessor", req: Req, input_token_logprobs: List ) -> None: """Process input token logprobs values and indices.""" - is_multi_item_scoring = self._is_multi_item_scoring(req) + is_multi_item_scoring = SchedulerOutputProcessorMixin._is_multi_item_scoring( + self, req + ) # Process logprob values - handle multi-item scoring vs regular requests if is_multi_item_scoring: @@ -741,12 +755,17 @@ class SchedulerOutputProcessorMixin: for x in input_token_logprobs_idx ] - def _process_input_top_logprobs(self: Scheduler, req: Req) -> None: + @staticmethod + def _process_input_top_logprobs( + self: "SchedulerLogprobResultProcessor", req: Req + ) -> None: """Process input top logprobs.""" if req.top_logprobs_num <= 0: return - is_multi_item_scoring = self._is_multi_item_scoring(req) + is_multi_item_scoring = SchedulerOutputProcessorMixin._is_multi_item_scoring( + self, req + ) # Initialize arrays - multi-item scoring starts empty, others start with None req.input_top_logprobs_val = [] if is_multi_item_scoring else [None] @@ -770,12 +789,17 @@ class SchedulerOutputProcessorMixin: req.temp_input_top_logprobs_idx = None req.temp_input_top_logprobs_val = None - def _process_input_token_ids_logprobs(self: Scheduler, req: Req) -> None: + @staticmethod + def _process_input_token_ids_logprobs( + self: "SchedulerLogprobResultProcessor", req: Req + ) -> None: """Process input token IDs logprobs.""" if req.token_ids_logprob is None: return - is_multi_item_scoring = self._is_multi_item_scoring(req) + is_multi_item_scoring = SchedulerOutputProcessorMixin._is_multi_item_scoring( + self, req + ) # Initialize arrays - multi-item scoring starts empty, others start with None req.input_token_ids_logprobs_val = [] if is_multi_item_scoring else [None] @@ -802,28 +826,39 @@ class SchedulerOutputProcessorMixin: req.temp_input_token_ids_logprobs_idx = None req.temp_input_token_ids_logprobs_val = None - def _calculate_relevant_tokens_len(self: Scheduler, req: Req) -> int: + @staticmethod + def _calculate_relevant_tokens_len( + self: "SchedulerLogprobResultProcessor", req: Req + ) -> int: """Calculate the expected length of logprob arrays based on whether multi-item scoring is enabled. For multi-item scoring, only delimiter positions have logprobs. For regular requests, all positions from logprob_start_len onwards have logprobs. """ - is_multi_item_scoring = self._is_multi_item_scoring(req) + is_multi_item_scoring = SchedulerOutputProcessorMixin._is_multi_item_scoring( + self, req + ) if is_multi_item_scoring: return len(req.multi_item_delimiter_indices) else: return len(req.origin_input_ids[req.logprob_start_len :]) - def _calculate_num_input_logprobs( - self: Scheduler, req: Req, extend_input_len: int, extend_logprob_start_len: int + @staticmethod + def calculate_num_input_logprobs( + self: "SchedulerLogprobResultProcessor", + req: Req, + extend_input_len: int, + extend_logprob_start_len: int, ) -> int: """Calculate the number of input logprobs based on whether multi-item scoring is enabled. For multi-item scoring, only delimiter positions have logprobs. For regular requests, all positions in the range have logprobs. """ - is_multi_item_scoring = self._is_multi_item_scoring(req) + is_multi_item_scoring = SchedulerOutputProcessorMixin._is_multi_item_scoring( + self, req + ) if is_multi_item_scoring: # Count pre-computed delimiter indices within the extend range @@ -836,7 +871,10 @@ class SchedulerOutputProcessorMixin: # Regular request: all tokens in the range return extend_input_len - extend_logprob_start_len - def _is_multi_item_scoring(self: Scheduler, req: Req) -> bool: + @staticmethod + def _is_multi_item_scoring( + self: "SchedulerLogprobResultProcessor", req: Req + ) -> bool: """Check if request uses multi-item scoring. Multi-item scoring applies to prefill-only requests when a delimiter @@ -849,8 +887,9 @@ class SchedulerOutputProcessorMixin: and req.multi_item_delimiter_indices is not None ) + @staticmethod def add_input_logprob_return_values( - self: Scheduler, + self: "SchedulerLogprobResultProcessor", i: int, req: Req, output: LogitsProcessorOutput, @@ -918,13 +957,19 @@ class SchedulerOutputProcessorMixin: assert req.input_top_logprobs_idx is None # Process all input logprob types using helper functions - self._process_input_token_logprobs(req, input_token_logprobs) - self._process_input_top_logprobs(req) + SchedulerOutputProcessorMixin._process_input_token_logprobs( + self, req, input_token_logprobs + ) + SchedulerOutputProcessorMixin._process_input_top_logprobs(self, req) - self._process_input_token_ids_logprobs(req) + SchedulerOutputProcessorMixin._process_input_token_ids_logprobs(self, req) if req.return_logprob: - relevant_tokens_len = self._calculate_relevant_tokens_len(req) + relevant_tokens_len = ( + SchedulerOutputProcessorMixin._calculate_relevant_tokens_len( + self, req + ) + ) assert len(req.input_token_logprobs_val) == relevant_tokens_len assert len(req.input_token_logprobs_idx) == relevant_tokens_len if req.top_logprobs_num > 0: @@ -934,8 +979,9 @@ class SchedulerOutputProcessorMixin: assert len(req.input_token_ids_logprobs_val) == relevant_tokens_len assert len(req.input_token_ids_logprobs_idx) == relevant_tokens_len + @staticmethod def add_logprob_return_values( - self: Scheduler, + self: "SchedulerLogprobResultProcessor", i: int, req: Req, pt: int, @@ -952,11 +998,13 @@ class SchedulerOutputProcessorMixin: # Note: For prefill-only requests with default logprob_start_len, this will be 0, # meaning we only compute output logprobs (which is the intended behavior) if num_input_logprobs > 0: - self.add_input_logprob_return_values( - i, req, output, pt, num_input_logprobs, last_prefill_chunk=True + SchedulerOutputProcessorMixin.add_input_logprob_return_values( + self, i, req, output, pt, num_input_logprobs, last_prefill_chunk=True ) else: - self._initialize_empty_logprob_containers(req) + SchedulerOutputProcessorMixin._initialize_empty_logprob_containers( + self, req + ) if req.top_logprobs_num > 0: req.output_top_logprobs_val.append(output.next_token_top_logprobs_val[i]) @@ -977,7 +1025,10 @@ class SchedulerOutputProcessorMixin: return num_input_logprobs - def _initialize_empty_logprob_containers(self: Scheduler, req: Req) -> None: + @staticmethod + def _initialize_empty_logprob_containers( + self: "SchedulerLogprobResultProcessor", req: Req + ) -> None: """ Initialize logprob fields to empty lists if unset.