Introduce SchedulerLogprobResultProcessor to own logprob state (#25632)
This commit is contained in:
@@ -536,6 +536,7 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
extend_input_len = extend_input_len_per_req[i]
|
extend_input_len = extend_input_len_per_req[i]
|
||||||
num_input_logprobs = extend_input_len - extend_logprob_start_len
|
num_input_logprobs = extend_input_len - extend_logprob_start_len
|
||||||
self.add_logprob_return_values(
|
self.add_logprob_return_values(
|
||||||
|
self.logprob_result_processor,
|
||||||
i,
|
i,
|
||||||
req,
|
req,
|
||||||
logprob_pt,
|
logprob_pt,
|
||||||
@@ -573,6 +574,7 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
# Update input logprobs.
|
# Update input logprobs.
|
||||||
num_input_logprobs = extend_input_len - extend_logprob_start_len
|
num_input_logprobs = extend_input_len - extend_logprob_start_len
|
||||||
self.add_input_logprob_return_values(
|
self.add_input_logprob_return_values(
|
||||||
|
self.logprob_result_processor,
|
||||||
i,
|
i,
|
||||||
req,
|
req,
|
||||||
logits_output,
|
logits_output,
|
||||||
|
|||||||
@@ -176,6 +176,9 @@ from sglang.srt.managers.scheduler_components.kv_events_publisher import (
|
|||||||
from sglang.srt.managers.scheduler_components.load_inquirer import (
|
from sglang.srt.managers.scheduler_components.load_inquirer import (
|
||||||
SchedulerLoadInquirer,
|
SchedulerLoadInquirer,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.managers.scheduler_components.logprob_result_processor import (
|
||||||
|
SchedulerLogprobResultProcessor,
|
||||||
|
)
|
||||||
from sglang.srt.managers.scheduler_components.metrics_reporter import (
|
from sglang.srt.managers.scheduler_components.metrics_reporter import (
|
||||||
RECORD_STEP_TIME,
|
RECORD_STEP_TIME,
|
||||||
PrefillStats,
|
PrefillStats,
|
||||||
@@ -734,6 +737,11 @@ class Scheduler(
|
|||||||
get_spec_total_num_forward_ct=lambda: self.metrics_reporter.spec_total_num_forward_ct,
|
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
|
self.is_initializing = False
|
||||||
|
|
||||||
def init_zbal_on_npu(self):
|
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,
|
ScheduleBatch,
|
||||||
Scheduler,
|
Scheduler,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.managers.scheduler_components.logprob_result_processor import (
|
||||||
|
SchedulerLogprobResultProcessor,
|
||||||
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -268,12 +271,16 @@ class SchedulerOutputProcessorMixin:
|
|||||||
extend_logprob_start_len = extend_logprob_start_len_per_req[i]
|
extend_logprob_start_len = extend_logprob_start_len_per_req[i]
|
||||||
extend_input_len = extend_input_len_per_req[i]
|
extend_input_len = extend_input_len_per_req[i]
|
||||||
|
|
||||||
num_input_logprobs = self._calculate_num_input_logprobs(
|
num_input_logprobs = self.calculate_num_input_logprobs(
|
||||||
req, extend_input_len, extend_logprob_start_len
|
self.logprob_result_processor,
|
||||||
|
req,
|
||||||
|
extend_input_len,
|
||||||
|
extend_logprob_start_len,
|
||||||
)
|
)
|
||||||
|
|
||||||
if req.return_logprob:
|
if req.return_logprob:
|
||||||
self.add_logprob_return_values(
|
self.add_logprob_return_values(
|
||||||
|
self.logprob_result_processor,
|
||||||
i,
|
i,
|
||||||
req,
|
req,
|
||||||
logprob_pt,
|
logprob_pt,
|
||||||
@@ -326,11 +333,15 @@ class SchedulerOutputProcessorMixin:
|
|||||||
extend_input_len = extend_input_len_per_req[i]
|
extend_input_len = extend_input_len_per_req[i]
|
||||||
if extend_logprob_start_len < extend_input_len:
|
if extend_logprob_start_len < extend_input_len:
|
||||||
# Update input logprobs.
|
# Update input logprobs.
|
||||||
num_input_logprobs = self._calculate_num_input_logprobs(
|
num_input_logprobs = self.calculate_num_input_logprobs(
|
||||||
req, extend_input_len, extend_logprob_start_len
|
self.logprob_result_processor,
|
||||||
|
req,
|
||||||
|
extend_input_len,
|
||||||
|
extend_logprob_start_len,
|
||||||
)
|
)
|
||||||
if req.return_logprob:
|
if req.return_logprob:
|
||||||
self.add_input_logprob_return_values(
|
self.add_input_logprob_return_values(
|
||||||
|
self.logprob_result_processor,
|
||||||
i,
|
i,
|
||||||
req,
|
req,
|
||||||
logits_output,
|
logits_output,
|
||||||
@@ -709,11 +720,14 @@ class SchedulerOutputProcessorMixin:
|
|||||||
actual_seq_len // mamba_track_interval * mamba_track_interval
|
actual_seq_len // mamba_track_interval * mamba_track_interval
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
def _process_input_token_logprobs(
|
def _process_input_token_logprobs(
|
||||||
self: Scheduler, req: Req, input_token_logprobs: List
|
self: "SchedulerLogprobResultProcessor", req: Req, input_token_logprobs: List
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Process input token logprobs values and indices."""
|
"""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
|
# Process logprob values - handle multi-item scoring vs regular requests
|
||||||
if is_multi_item_scoring:
|
if is_multi_item_scoring:
|
||||||
@@ -741,12 +755,17 @@ class SchedulerOutputProcessorMixin:
|
|||||||
for x in input_token_logprobs_idx
|
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."""
|
"""Process input top logprobs."""
|
||||||
if req.top_logprobs_num <= 0:
|
if req.top_logprobs_num <= 0:
|
||||||
return
|
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
|
# Initialize arrays - multi-item scoring starts empty, others start with None
|
||||||
req.input_top_logprobs_val = [] if is_multi_item_scoring else [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_idx = None
|
||||||
req.temp_input_top_logprobs_val = 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."""
|
"""Process input token IDs logprobs."""
|
||||||
if req.token_ids_logprob is None:
|
if req.token_ids_logprob is None:
|
||||||
return
|
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
|
# Initialize arrays - multi-item scoring starts empty, others start with None
|
||||||
req.input_token_ids_logprobs_val = [] if is_multi_item_scoring else [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_idx = None
|
||||||
req.temp_input_token_ids_logprobs_val = 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.
|
"""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 multi-item scoring, only delimiter positions have logprobs.
|
||||||
For regular requests, all positions from logprob_start_len onwards 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:
|
if is_multi_item_scoring:
|
||||||
return len(req.multi_item_delimiter_indices)
|
return len(req.multi_item_delimiter_indices)
|
||||||
else:
|
else:
|
||||||
return len(req.origin_input_ids[req.logprob_start_len :])
|
return len(req.origin_input_ids[req.logprob_start_len :])
|
||||||
|
|
||||||
def _calculate_num_input_logprobs(
|
@staticmethod
|
||||||
self: Scheduler, req: Req, extend_input_len: int, extend_logprob_start_len: int
|
def calculate_num_input_logprobs(
|
||||||
|
self: "SchedulerLogprobResultProcessor",
|
||||||
|
req: Req,
|
||||||
|
extend_input_len: int,
|
||||||
|
extend_logprob_start_len: int,
|
||||||
) -> int:
|
) -> int:
|
||||||
"""Calculate the number of input logprobs based on whether multi-item scoring is enabled.
|
"""Calculate the number of input logprobs based on whether multi-item scoring is enabled.
|
||||||
|
|
||||||
For multi-item scoring, only delimiter positions have logprobs.
|
For multi-item scoring, only delimiter positions have logprobs.
|
||||||
For regular requests, all positions in the range 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:
|
if is_multi_item_scoring:
|
||||||
# Count pre-computed delimiter indices within the extend range
|
# Count pre-computed delimiter indices within the extend range
|
||||||
@@ -836,7 +871,10 @@ class SchedulerOutputProcessorMixin:
|
|||||||
# Regular request: all tokens in the range
|
# Regular request: all tokens in the range
|
||||||
return extend_input_len - extend_logprob_start_len
|
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.
|
"""Check if request uses multi-item scoring.
|
||||||
|
|
||||||
Multi-item scoring applies to prefill-only requests when a delimiter
|
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
|
and req.multi_item_delimiter_indices is not None
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
def add_input_logprob_return_values(
|
def add_input_logprob_return_values(
|
||||||
self: Scheduler,
|
self: "SchedulerLogprobResultProcessor",
|
||||||
i: int,
|
i: int,
|
||||||
req: Req,
|
req: Req,
|
||||||
output: LogitsProcessorOutput,
|
output: LogitsProcessorOutput,
|
||||||
@@ -918,13 +957,19 @@ class SchedulerOutputProcessorMixin:
|
|||||||
assert req.input_top_logprobs_idx is None
|
assert req.input_top_logprobs_idx is None
|
||||||
|
|
||||||
# Process all input logprob types using helper functions
|
# Process all input logprob types using helper functions
|
||||||
self._process_input_token_logprobs(req, input_token_logprobs)
|
SchedulerOutputProcessorMixin._process_input_token_logprobs(
|
||||||
self._process_input_top_logprobs(req)
|
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:
|
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_val) == relevant_tokens_len
|
||||||
assert len(req.input_token_logprobs_idx) == relevant_tokens_len
|
assert len(req.input_token_logprobs_idx) == relevant_tokens_len
|
||||||
if req.top_logprobs_num > 0:
|
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_val) == relevant_tokens_len
|
||||||
assert len(req.input_token_ids_logprobs_idx) == relevant_tokens_len
|
assert len(req.input_token_ids_logprobs_idx) == relevant_tokens_len
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
def add_logprob_return_values(
|
def add_logprob_return_values(
|
||||||
self: Scheduler,
|
self: "SchedulerLogprobResultProcessor",
|
||||||
i: int,
|
i: int,
|
||||||
req: Req,
|
req: Req,
|
||||||
pt: int,
|
pt: int,
|
||||||
@@ -952,11 +998,13 @@ class SchedulerOutputProcessorMixin:
|
|||||||
# Note: For prefill-only requests with default logprob_start_len, this will be 0,
|
# 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)
|
# meaning we only compute output logprobs (which is the intended behavior)
|
||||||
if num_input_logprobs > 0:
|
if num_input_logprobs > 0:
|
||||||
self.add_input_logprob_return_values(
|
SchedulerOutputProcessorMixin.add_input_logprob_return_values(
|
||||||
i, req, output, pt, num_input_logprobs, last_prefill_chunk=True
|
self, i, req, output, pt, num_input_logprobs, last_prefill_chunk=True
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self._initialize_empty_logprob_containers(req)
|
SchedulerOutputProcessorMixin._initialize_empty_logprob_containers(
|
||||||
|
self, req
|
||||||
|
)
|
||||||
|
|
||||||
if req.top_logprobs_num > 0:
|
if req.top_logprobs_num > 0:
|
||||||
req.output_top_logprobs_val.append(output.next_token_top_logprobs_val[i])
|
req.output_top_logprobs_val.append(output.next_token_top_logprobs_val[i])
|
||||||
@@ -977,7 +1025,10 @@ class SchedulerOutputProcessorMixin:
|
|||||||
|
|
||||||
return num_input_logprobs
|
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.
|
Initialize logprob fields to empty lists if unset.
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user