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