Introduce SchedulerLogprobResultProcessor to own logprob state (#25632)

This commit is contained in:
fzyzcjy
2026-05-18 18:42:59 +08:00
committed by GitHub
parent cf12070a0f
commit e737f61b29
4 changed files with 100 additions and 26 deletions
@@ -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,
+8
View File
@@ -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):
@@ -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
@@ -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.