[perf] Assemble flat prompt top logprobs scheduler-side as numpy arrays (#32223)

This commit is contained in:
Sam Shleifer
2026-07-31 18:00:51 -07:00
committed by GitHub
parent ca07917c58
commit 58974ca16c
9 changed files with 593 additions and 22 deletions
@@ -462,6 +462,9 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
output_token_logprobs_idx=recv_obj.output_token_logprobs_idx, output_token_logprobs_idx=recv_obj.output_token_logprobs_idx,
input_top_logprobs_val=recv_obj.input_top_logprobs_val, input_top_logprobs_val=recv_obj.input_top_logprobs_val,
input_top_logprobs_idx=recv_obj.input_top_logprobs_idx, input_top_logprobs_idx=recv_obj.input_top_logprobs_idx,
input_top_logprobs_val_flat=recv_obj.input_top_logprobs_val_flat,
input_top_logprobs_idx_flat=recv_obj.input_top_logprobs_idx_flat,
input_top_logprobs_flat_null_prefix=recv_obj.input_top_logprobs_flat_null_prefix,
output_top_logprobs_val=recv_obj.output_top_logprobs_val, output_top_logprobs_val=recv_obj.output_top_logprobs_val,
output_top_logprobs_idx=recv_obj.output_top_logprobs_idx, output_top_logprobs_idx=recv_obj.output_top_logprobs_idx,
input_token_ids_logprobs_val=recv_obj.input_token_ids_logprobs_val, input_token_ids_logprobs_val=recv_obj.input_token_ids_logprobs_val,
+53
View File
@@ -38,6 +38,7 @@ from typing import (
List, List,
Literal, Literal,
Optional, Optional,
Tuple,
Type, Type,
Union, Union,
) )
@@ -831,6 +832,10 @@ class TokenizedGenerateReqInput(BaseReq, kw_only=True):
stream: bool stream: bool
# Whether to return sparse output-token support from top-k/top-p/min-p sampling. # Whether to return sparse output-token support from top-k/top-p/min-p sampling.
return_sampling_mask: bool = False return_sampling_mask: bool = False
# Assemble prompt top logprobs as flat arrays scheduler-side (see
# GenerateReqInput.return_flat_raw_top_logprobs). The b64 flag stays
# tokenizer-manager-side: the scheduler ships arrays either way.
return_flat_raw_top_logprobs: bool = False
# Whether to return hidden states # Whether to return hidden states
return_hidden_states: bool = False return_hidden_states: bool = False
@@ -1229,6 +1234,39 @@ CachedTokensDetails = Dict[str, Union[int, str]]
FinishReasonDict = Dict[str, Optional[Union[str, int, List[int]]]] FinishReasonDict = Dict[str, Optional[Union[str, int, List[int]]]]
def build_flat_input_top_logprobs_arrays(
input_top_logprobs_val: List[Optional[List[float]]],
input_top_logprobs_idx: List[Optional[List[int]]],
top_logprobs_num: int,
) -> Tuple[np.ndarray, np.ndarray, int]:
"""Convert nested per-position prompt top logprob rows into the flat
arrays of the `return_flat_raw_top_logprobs` response format.
Returns (float32 values [rows, k], int32 token ids [rows, k],
null_prefix). The leading null rows are counted into null_prefix and
excluded from the arrays. Raises ValueError when the rows are not
representable by (shape, null_prefix): interior nulls or ragged k,
e.g. multi-item scoring.
"""
num_rows = len(input_top_logprobs_val)
null_prefix = 0
while null_prefix < num_rows and not input_top_logprobs_val[null_prefix]:
null_prefix += 1
val_rows = input_top_logprobs_val[null_prefix:]
idx_rows = input_top_logprobs_idx[null_prefix:]
k = len(val_rows[0]) if val_rows else top_logprobs_num
for offset, row in enumerate(val_rows):
if row is None or len(row) != k:
raise ValueError(
"return_flat_raw_top_logprobs requires rectangular top logprob "
f"rows with nulls only in the leading prefix; row {null_prefix + offset} "
f"has {None if row is None else len(row)} entries (expected {k})."
)
val_arr = np.asarray(val_rows, dtype=np.float32).reshape(len(val_rows), k)
idx_arr = np.asarray(idx_rows, dtype=np.int32).reshape(len(idx_rows), k)
return val_arr, idx_arr, null_prefix
class BatchTokenIDOutput(BaseBatchReq, kw_only=True): class BatchTokenIDOutput(BaseBatchReq, kw_only=True):
# The finish reason # The finish reason
finished_reasons: List[Optional[FinishReasonDict]] finished_reasons: List[Optional[FinishReasonDict]]
@@ -1319,6 +1357,15 @@ class BatchTokenIDOutput(BaseBatchReq, kw_only=True):
spec_correct_drafts_histogram: Optional[List[List[int]]] = None spec_correct_drafts_histogram: Optional[List[List[int]]] = None
spec_cap_lens_histogram: Optional[List[List[int]]] = None spec_cap_lens_histogram: Optional[List[List[int]]] = None
# Scheduler-side flat assembly of prompt top logprobs for requests with
# return_flat_raw_top_logprobs: float32 / int32 [rows, k] arrays plus the
# leading-null count (see build_flat_input_top_logprobs_arrays). For such
# requests the nested input_top_logprobs_val/idx entry is empty. None when
# no request in the batch uses the flat format.
input_top_logprobs_val_flat: Optional[List[Optional[np.ndarray]]] = None
input_top_logprobs_idx_flat: Optional[List[Optional[np.ndarray]]] = None
input_top_logprobs_flat_null_prefix: Optional[List[Optional[int]]] = None
class BatchStrOutput(BaseBatchReq, kw_only=True): class BatchStrOutput(BaseBatchReq, kw_only=True):
# The finish reason # The finish reason
@@ -1401,6 +1448,12 @@ class BatchStrOutput(BaseBatchReq, kw_only=True):
spec_correct_drafts_histogram: Optional[List[List[int]]] = None spec_correct_drafts_histogram: Optional[List[List[int]]] = None
spec_cap_lens_histogram: Optional[List[List[int]]] = None spec_cap_lens_histogram: Optional[List[List[int]]] = None
# Detokenizer pass-through for the scheduler-side flat prompt top logprob
# arrays; see BatchTokenIDOutput.input_top_logprobs_val_flat.
input_top_logprobs_val_flat: Optional[List[Optional[np.ndarray]]] = None
input_top_logprobs_idx_flat: Optional[List[Optional[np.ndarray]]] = None
input_top_logprobs_flat_null_prefix: Optional[List[Optional[int]]] = None
class BatchEmbeddingOutput(BaseBatchReq, kw_only=True): class BatchEmbeddingOutput(BaseBatchReq, kw_only=True):
# The finish reason # The finish reason
@@ -211,6 +211,15 @@ def _handle_output_by_index(output, i):
input_top_logprobs_idx=_extract_field_by_index( input_top_logprobs_idx=_extract_field_by_index(
output, "input_top_logprobs_idx", i, check_length=False output, "input_top_logprobs_idx", i, check_length=False
), ),
input_top_logprobs_val_flat=_extract_field_by_index(
output, "input_top_logprobs_val_flat", i, check_length=False
),
input_top_logprobs_idx_flat=_extract_field_by_index(
output, "input_top_logprobs_idx_flat", i, check_length=False
),
input_top_logprobs_flat_null_prefix=_extract_field_by_index(
output, "input_top_logprobs_flat_null_prefix", i, check_length=False
),
output_top_logprobs_val=_extract_field_by_index( output_top_logprobs_val=_extract_field_by_index(
output, "output_top_logprobs_val", i, check_length=False output, "output_top_logprobs_val", i, check_length=False
), ),
@@ -319,6 +328,15 @@ def _handle_output_by_index(output, i):
input_top_logprobs_idx=_extract_field_by_index( input_top_logprobs_idx=_extract_field_by_index(
output, "input_top_logprobs_idx", i, check_length=False output, "input_top_logprobs_idx", i, check_length=False
), ),
input_top_logprobs_val_flat=_extract_field_by_index(
output, "input_top_logprobs_val_flat", i, check_length=False
),
input_top_logprobs_idx_flat=_extract_field_by_index(
output, "input_top_logprobs_idx_flat", i, check_length=False
),
input_top_logprobs_flat_null_prefix=_extract_field_by_index(
output, "input_top_logprobs_flat_null_prefix", i, check_length=False
),
output_top_logprobs_val=_extract_field_by_index( output_top_logprobs_val=_extract_field_by_index(
output, "output_top_logprobs_val", i, check_length=False output, "output_top_logprobs_val", i, check_length=False
), ),
@@ -687,6 +687,12 @@ class ReqLogprob:
input_token_logprobs_idx: Optional[List[int]] = None input_token_logprobs_idx: Optional[List[int]] = None
input_top_logprobs_val: Optional[List[List[float]]] = None input_top_logprobs_val: Optional[List[List[float]]] = None
input_top_logprobs_idx: Optional[List[List[int]]] = None input_top_logprobs_idx: Optional[List[List[int]]] = None
# Flat replacements for the rows above (see
# build_flat_input_top_logprobs_arrays); when set, the nested rows are
# emptied and the arrays ship instead.
input_top_logprobs_val_flat: Optional[np.ndarray] = None
input_top_logprobs_idx_flat: Optional[np.ndarray] = None
input_top_logprobs_flat_null_prefix: Optional[int] = None
input_token_ids_logprobs_val: Optional[List[List[float]]] = None input_token_ids_logprobs_val: Optional[List[List[float]]] = None
input_token_ids_logprobs_idx: Optional[List[List[int]]] = None input_token_ids_logprobs_idx: Optional[List[List[int]]] = None
output_token_logprobs_val: Optional[list] = None output_token_logprobs_val: Optional[list] = None
@@ -725,6 +731,7 @@ class Req(ReqDllmMixin):
dllm_config: Optional[DllmConfig] = None, dllm_config: Optional[DllmConfig] = None,
token_ids_logprob: List[int] = None, token_ids_logprob: List[int] = None,
return_sampling_mask: bool = False, return_sampling_mask: bool = False,
return_flat_raw_top_logprobs: bool = False,
stream: bool = False, stream: bool = False,
origin_input_ids_unpadded: Optional[array[int]] = None, origin_input_ids_unpadded: Optional[array[int]] = None,
lora_id: Optional[str] = None, lora_id: Optional[str] = None,
@@ -943,6 +950,7 @@ class Req(ReqDllmMixin):
self.temp_scaled_logprobs = False self.temp_scaled_logprobs = False
self.top_p_normalized_logprobs = False self.top_p_normalized_logprobs = False
self.return_sampling_mask = return_sampling_mask self.return_sampling_mask = return_sampling_mask
self.return_flat_raw_top_logprobs = return_flat_raw_top_logprobs
# Logprobs (return values) # Logprobs (return values)
# True means the input logprob has been already sent to detokenizer. # True means the input logprob has been already sent to detokenizer.
+1
View File
@@ -2242,6 +2242,7 @@ class Scheduler(
top_logprobs_num=recv_req.top_logprobs_num, top_logprobs_num=recv_req.top_logprobs_num,
token_ids_logprob=recv_req.token_ids_logprob, token_ids_logprob=recv_req.token_ids_logprob,
return_sampling_mask=recv_req.return_sampling_mask, return_sampling_mask=recv_req.return_sampling_mask,
return_flat_raw_top_logprobs=recv_req.return_flat_raw_top_logprobs,
stream=recv_req.stream, stream=recv_req.stream,
lora_id=recv_req.lora_id, lora_id=recv_req.lora_id,
session_id=recv_req.session_id, session_id=recv_req.session_id,
@@ -1,5 +1,6 @@
from __future__ import annotations from __future__ import annotations
import logging
from dataclasses import dataclass from dataclasses import dataclass
from typing import ( from typing import (
List, List,
@@ -10,6 +11,7 @@ import torch
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.managers.io_struct import build_flat_input_top_logprobs_arrays
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
from sglang.srt.runtime_context import get_exec from sglang.srt.runtime_context import get_exec
from sglang.srt.server_args import ( from sglang.srt.server_args import (
@@ -17,6 +19,8 @@ from sglang.srt.server_args import (
ServerArgs, ServerArgs,
) )
logger = logging.getLogger(__name__)
@dataclass(kw_only=True, slots=True, frozen=True) @dataclass(kw_only=True, slots=True, frozen=True)
class SchedulerLogprobResultProcessor: class SchedulerLogprobResultProcessor:
@@ -84,6 +88,35 @@ class SchedulerLogprobResultProcessor:
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 _flatten_input_top_logprobs(self, req: Req) -> None:
"""Replace the nested input top logprob rows with flat arrays for
requests that opted into return_flat_raw_top_logprobs, so the batch
output ships two ndarrays instead of num_positions * k python lists.
"""
if req.logprob.top_logprobs_num <= 0:
return
try:
(
req.logprob.input_top_logprobs_val_flat,
req.logprob.input_top_logprobs_idx_flat,
req.logprob.input_top_logprobs_flat_null_prefix,
) = build_flat_input_top_logprobs_arrays(
req.logprob.input_top_logprobs_val,
req.logprob.input_top_logprobs_idx,
req.logprob.top_logprobs_num,
)
except ValueError as e:
# Unrepresentable rows (e.g. multi-item scoring): keep the nested
# format, mirroring the tokenizer manager fallback.
logger.warning(
"Falling back to nested input top logprobs for rid=%s: %s",
req.rid,
e,
)
return
req.logprob.input_top_logprobs_val = []
req.logprob.input_top_logprobs_idx = []
def _process_input_token_ids_logprobs(self, req: Req) -> None: def _process_input_token_ids_logprobs(self, req: Req) -> None:
"""Process input token IDs logprobs.""" """Process input token IDs logprobs."""
if req.logprob.token_ids_logprob is None: if req.logprob.token_ids_logprob is None:
@@ -265,6 +298,10 @@ class SchedulerLogprobResultProcessor:
== relevant_tokens_len == relevant_tokens_len
) )
# After the length checks: the flat arrays replace the nested rows.
if req.return_flat_raw_top_logprobs:
self._flatten_input_top_logprobs(req)
def add_logprob_return_values( def add_logprob_return_values(
self, self,
i: int, i: int,
@@ -312,6 +312,12 @@ class _GenerationStreamAccumulator:
output_token_logprobs_idx: Optional[list] = None output_token_logprobs_idx: Optional[list] = None
input_top_logprobs_val: Optional[list] = None input_top_logprobs_val: Optional[list] = None
input_top_logprobs_idx: Optional[list] = None input_top_logprobs_idx: Optional[list] = None
# Per-request flat prompt top logprob arrays (return_flat_raw_top_logprobs);
# None entries for requests on the nested format.
input_top_logprobs_val_flat: Optional[list] = None
input_top_logprobs_idx_flat: Optional[list] = None
input_top_logprobs_flat_null_prefix: Optional[list] = None
has_input_top_logprobs_flat: bool = False
output_top_logprobs_val: Optional[list] = None output_top_logprobs_val: Optional[list] = None
output_top_logprobs_idx: Optional[list] = None output_top_logprobs_idx: Optional[list] = None
input_token_ids_logprobs_val: Optional[list] = None input_token_ids_logprobs_val: Optional[list] = None
@@ -340,6 +346,9 @@ class _GenerationStreamAccumulator:
self.output_token_logprobs_idx = [] self.output_token_logprobs_idx = []
self.input_top_logprobs_val = [] self.input_top_logprobs_val = []
self.input_top_logprobs_idx = [] self.input_top_logprobs_idx = []
self.input_top_logprobs_val_flat = []
self.input_top_logprobs_idx_flat = []
self.input_top_logprobs_flat_null_prefix = []
self.output_top_logprobs_val = [] self.output_top_logprobs_val = []
self.output_top_logprobs_idx = [] self.output_top_logprobs_idx = []
self.input_token_ids_logprobs_val = [] self.input_token_ids_logprobs_val = []
@@ -464,6 +473,17 @@ class _GenerationStreamAccumulator:
) )
self.input_top_logprobs_val.append(req.logprob.input_top_logprobs_val) self.input_top_logprobs_val.append(req.logprob.input_top_logprobs_val)
self.input_top_logprobs_idx.append(req.logprob.input_top_logprobs_idx) self.input_top_logprobs_idx.append(req.logprob.input_top_logprobs_idx)
self.input_top_logprobs_val_flat.append(
req.logprob.input_top_logprobs_val_flat
)
self.input_top_logprobs_idx_flat.append(
req.logprob.input_top_logprobs_idx_flat
)
self.input_top_logprobs_flat_null_prefix.append(
req.logprob.input_top_logprobs_flat_null_prefix
)
if req.logprob.input_top_logprobs_val_flat is not None:
self.has_input_top_logprobs_flat = True
self.input_token_ids_logprobs_val.append( self.input_token_ids_logprobs_val.append(
req.logprob.input_token_ids_logprobs_val req.logprob.input_token_ids_logprobs_val
) )
@@ -476,6 +496,9 @@ class _GenerationStreamAccumulator:
self.input_token_logprobs_idx.append([]) self.input_token_logprobs_idx.append([])
self.input_top_logprobs_val.append([]) self.input_top_logprobs_val.append([])
self.input_top_logprobs_idx.append([]) self.input_top_logprobs_idx.append([])
self.input_top_logprobs_val_flat.append(None)
self.input_top_logprobs_idx_flat.append(None)
self.input_top_logprobs_flat_null_prefix.append(None)
self.input_token_ids_logprobs_val.append([]) self.input_token_ids_logprobs_val.append([])
self.input_token_ids_logprobs_idx.append([]) self.input_token_ids_logprobs_idx.append([])
@@ -613,6 +636,23 @@ class _GenerationStreamAccumulator:
output_token_logprobs_idx=self.output_token_logprobs_idx, output_token_logprobs_idx=self.output_token_logprobs_idx,
input_top_logprobs_val=self.input_top_logprobs_val, input_top_logprobs_val=self.input_top_logprobs_val,
input_top_logprobs_idx=self.input_top_logprobs_idx, input_top_logprobs_idx=self.input_top_logprobs_idx,
# None on the common path so the wire payload is unchanged when no
# request in the batch uses the flat format.
input_top_logprobs_val_flat=(
self.input_top_logprobs_val_flat
if self.has_input_top_logprobs_flat
else None
),
input_top_logprobs_idx_flat=(
self.input_top_logprobs_idx_flat
if self.has_input_top_logprobs_flat
else None
),
input_top_logprobs_flat_null_prefix=(
self.input_top_logprobs_flat_null_prefix
if self.has_input_top_logprobs_flat
else None
),
output_top_logprobs_val=self.output_top_logprobs_val, output_top_logprobs_val=self.output_top_logprobs_val,
output_top_logprobs_idx=self.output_top_logprobs_idx, output_top_logprobs_idx=self.output_top_logprobs_idx,
input_token_ids_logprobs_val=self.input_token_ids_logprobs_val, input_token_ids_logprobs_val=self.input_token_ids_logprobs_val,
+65 -20
View File
@@ -84,6 +84,7 @@ from sglang.srt.managers.io_struct import (
UpdateWeightFromDiskReqOutput, UpdateWeightFromDiskReqOutput,
async_sock_recv, async_sock_recv,
async_sock_send, async_sock_send,
build_flat_input_top_logprobs_arrays,
sock_send, sock_send,
unwrap_from_pickle, unwrap_from_pickle,
) )
@@ -233,6 +234,12 @@ class ReqState:
# prefill chunks arrive, so streaming decode chunks reuse the payload. # prefill chunks arrive, so streaming decode chunks reuse the payload.
input_top_logprobs_flat_fields: Optional[Dict[str, Any]] = None input_top_logprobs_flat_fields: Optional[Dict[str, Any]] = None
input_top_logprobs_flat_num_rows: int = -1 input_top_logprobs_flat_num_rows: int = -1
# Scheduler-assembled flat arrays (val float32 [rows, k], idx int32
# [rows, k], null_prefix), sent once at prefill completion. When present,
# the nested input_top_logprobs_val/idx above stay empty.
input_top_logprobs_scheduler_flat: Optional[Tuple[np.ndarray, np.ndarray, int]] = (
None
)
# For detokenized logprobs # For detokenized logprobs
input_token_logprobs: List[Any] = dataclasses.field(default_factory=list) input_token_logprobs: List[Any] = dataclasses.field(default_factory=list)
@@ -278,26 +285,38 @@ def _build_flat_input_top_logprobs_fields(
base64 contiguous little-endian binary; the dtype marker fields let the base64 contiguous little-endian binary; the dtype marker fields let the
widths change later without a wire break. widths change later without a wire break.
""" """
num_rows = len(input_top_logprobs_val) val_arr, idx_arr, null_prefix = build_flat_input_top_logprobs_arrays(
null_prefix = 0 input_top_logprobs_val, input_top_logprobs_idx, top_logprobs_num
while null_prefix < num_rows and not input_top_logprobs_val[null_prefix]:
null_prefix += 1
val_rows = input_top_logprobs_val[null_prefix:]
idx_rows = input_top_logprobs_idx[null_prefix:]
k = len(val_rows[0]) if val_rows else top_logprobs_num
for offset, row in enumerate(val_rows):
if row is None or len(row) != k:
# Not representable by (shape, null_prefix); e.g. multi-item scoring.
raise ValueError(
"return_flat_raw_top_logprobs requires rectangular top logprob "
f"rows with nulls only in the leading prefix; row {null_prefix + offset} "
f"has {None if row is None else len(row)} entries (expected {k})."
) )
if return_b64:
return _build_flat_input_top_logprobs_fields_from_arrays(
val_arr, idx_arr, null_prefix, return_b64=True
)
# Flatten the original python rows so the JSON numbers keep their full
# (float64) precision, matching the pre-scheduler-flat output.
return {
"input_top_logprobs_val_flat": [
v for row in input_top_logprobs_val[null_prefix:] for v in row
],
"input_top_logprobs_idx_flat": [
i for row in input_top_logprobs_idx[null_prefix:] for i in row
],
"input_top_logprobs_shape": [val_arr.shape[0], val_arr.shape[1]],
"input_top_logprobs_null_prefix": null_prefix,
}
def _build_flat_input_top_logprobs_fields_from_arrays(
val_arr: np.ndarray,
idx_arr: np.ndarray,
null_prefix: int,
return_b64: bool = False,
) -> Dict[str, Any]:
"""Build the flat response fields from scheduler-assembled [rows, k]
arrays (see `_build_flat_input_top_logprobs_fields` for the field
semantics)."""
fields: Dict[str, Any] = {} fields: Dict[str, Any] = {}
if return_b64: if return_b64:
val_arr = np.asarray(val_rows, dtype=np.float32)
idx_arr = np.asarray(idx_rows, dtype=np.int32)
fields["input_top_logprobs_val_flat_b64"] = pybase64.b64encode( fields["input_top_logprobs_val_flat_b64"] = pybase64.b64encode(
val_arr.tobytes() val_arr.tobytes()
).decode("utf-8") ).decode("utf-8")
@@ -307,9 +326,9 @@ def _build_flat_input_top_logprobs_fields(
fields["input_top_logprobs_val_flat_b64_dtype"] = "float32" fields["input_top_logprobs_val_flat_b64_dtype"] = "float32"
fields["input_top_logprobs_idx_flat_b64_dtype"] = "int32" fields["input_top_logprobs_idx_flat_b64_dtype"] = "int32"
else: else:
fields["input_top_logprobs_val_flat"] = [v for row in val_rows for v in row] fields["input_top_logprobs_val_flat"] = val_arr.reshape(-1).tolist()
fields["input_top_logprobs_idx_flat"] = [i for row in idx_rows for i in row] fields["input_top_logprobs_idx_flat"] = idx_arr.reshape(-1).tolist()
fields["input_top_logprobs_shape"] = [len(val_rows), k] fields["input_top_logprobs_shape"] = [val_arr.shape[0], val_arr.shape[1]]
fields["input_top_logprobs_null_prefix"] = null_prefix fields["input_top_logprobs_null_prefix"] = null_prefix
return fields return fields
@@ -1264,6 +1283,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
top_logprobs_num=obj.top_logprobs_num, top_logprobs_num=obj.top_logprobs_num,
token_ids_logprob=obj.token_ids_logprob, token_ids_logprob=obj.token_ids_logprob,
return_sampling_mask=obj.return_sampling_mask, return_sampling_mask=obj.return_sampling_mask,
return_flat_raw_top_logprobs=obj.return_flat_raw_top_logprobs,
stream=obj.stream, stream=obj.stream,
rid=obj.rid, rid=obj.rid,
http_worker_ipc=obj.http_worker_ipc, http_worker_ipc=obj.http_worker_ipc,
@@ -2297,7 +2317,23 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Guarded by the caller's return_logprob check, so obj is a # Guarded by the caller's return_logprob check, so obj is a
# GenerateReqInput here. # GenerateReqInput here.
use_flat = state.obj.return_flat_raw_top_logprobs use_flat = state.obj.return_flat_raw_top_logprobs
if use_flat: if use_flat and state.input_top_logprobs_scheduler_flat is not None:
# The scheduler already assembled the flat arrays (sent once
# at prefill completion); encode them directly.
if state.input_top_logprobs_flat_fields is None:
val_arr, idx_arr, null_prefix = (
state.input_top_logprobs_scheduler_flat
)
state.input_top_logprobs_flat_fields = (
_build_flat_input_top_logprobs_fields_from_arrays(
val_arr,
idx_arr,
null_prefix,
return_b64=state.obj.return_flat_raw_top_logprobs_b64,
)
)
meta_info.update(state.input_top_logprobs_flat_fields)
elif use_flat:
# Flat replaces nested for the input side only. # Flat replaces nested for the input side only.
if state.input_top_logprobs_flat_num_rows != len( if state.input_top_logprobs_flat_num_rows != len(
state.input_top_logprobs_val state.input_top_logprobs_val
@@ -2424,6 +2460,15 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
state.input_top_logprobs_idx.extend( state.input_top_logprobs_idx.extend(
recv_obj.input_top_logprobs_idx[recv_obj_index] recv_obj.input_top_logprobs_idx[recv_obj_index]
) )
if (
recv_obj.input_top_logprobs_val_flat is not None
and recv_obj.input_top_logprobs_val_flat[recv_obj_index] is not None
):
state.input_top_logprobs_scheduler_flat = (
recv_obj.input_top_logprobs_val_flat[recv_obj_index],
recv_obj.input_top_logprobs_idx_flat[recv_obj_index],
recv_obj.input_top_logprobs_flat_null_prefix[recv_obj_index],
)
state.output_top_logprobs_val.extend( state.output_top_logprobs_val.extend(
recv_obj.output_top_logprobs_val[recv_obj_index] recv_obj.output_top_logprobs_val[recv_obj_index]
) )
@@ -6,8 +6,11 @@ import asyncio
import base64 import base64
import json import json
import os import os
import pickle
import time import time
import unittest import unittest
from array import array
from types import SimpleNamespace
import numpy as np import numpy as np
@@ -16,13 +19,25 @@ from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
maybe_stub_sgl_kernel() maybe_stub_sgl_kernel()
from sglang.srt.managers.io_struct import GenerateReqInput from sglang.srt.managers.io_struct import (
BatchTokenIDOutput,
GenerateReqInput,
build_flat_input_top_logprobs_arrays,
msgpack_decode,
msgpack_encode,
)
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.managers.scheduler_components.logprob_result_processor import (
SchedulerLogprobResultProcessor,
)
from sglang.srt.managers.tokenizer_manager import ( from sglang.srt.managers.tokenizer_manager import (
ReqState, ReqState,
TokenizerManager, TokenizerManager,
_build_flat_input_top_logprobs_fields, _build_flat_input_top_logprobs_fields,
_build_flat_input_top_logprobs_fields_from_arrays,
) )
from sglang.srt.observability.req_time_stats import APIServerReqTimeStats from sglang.srt.observability.req_time_stats import APIServerReqTimeStats
from sglang.srt.sampling.sampling_params import SamplingParams
register_cpu_ci(est_time=10, suite="base-a-test-cpu") register_cpu_ci(est_time=10, suite="base-a-test-cpu")
@@ -31,6 +46,12 @@ register_cpu_ci(est_time=10, suite="base-a-test-cpu")
_VAL_ROWS = [None, [-0.1, -2.5], [-0.3, -1.5], [-0.05, -4.0]] _VAL_ROWS = [None, [-0.1, -2.5], [-0.3, -1.5], [-0.05, -4.0]]
_IDX_ROWS = [None, [11, 22], [33, 44], [55, 66]] _IDX_ROWS = [None, [11, 22], [33, 44], [55, 66]]
# Float32-exact values for scheduler-flat equivalence tests: the scheduler
# ships float32 arrays, so equality against the python-float rows needs values
# that survive the float64 -> float32 round trip (production logprobs do, being
# computed in float32).
_EXACT_VAL_ROWS = [None, [-0.5, -2.5], [-0.25, -1.5], [-0.125, -4.0]]
class _TokenizerManagerStub: class _TokenizerManagerStub:
"""Borrow the real logprob meta_info methods without a full manager.""" """Borrow the real logprob meta_info methods without a full manager."""
@@ -296,6 +317,316 @@ class TestB64MetaInfo(CustomTestCase):
) )
def _make_logprob_processor() -> SchedulerLogprobResultProcessor:
# The processor only reads enable_mis and vocab_size from these.
return SchedulerLogprobResultProcessor(
server_args=SimpleNamespace(enable_mis=False),
model_config=SimpleNamespace(vocab_size=1_000_000),
)
# Per-position rows as computed during prefill: one row per prompt position
# from logprob_start_len on, the last row being the sampling position that
# scheduler-side assembly pops.
_SCHED_VAL_ROWS = [
[-0.5, -2.5],
[-0.25, -1.5],
[-0.125, -4.0],
[-1.0, -3.0],
[-2.0, -5.0],
]
_SCHED_IDX_ROWS = [[11, 22], [33, 44], [55, 66], [77, 88], [99, 100]]
class TestSchedulerFlatAssembly(CustomTestCase):
"""Scheduler-side flat assembly in the logprob result processor."""
def _make_req(self, flat: bool, num_tokens: int = 5) -> Req:
return Req(
"r0",
"",
array("q", range(1, num_tokens + 1)),
SamplingParams(),
return_logprob=True,
top_logprobs_num=2,
return_flat_raw_top_logprobs=flat,
)
def _run_prefill(self, req: Req, chunk_sizes) -> None:
processor = _make_logprob_processor()
token_logprobs = [row[0] for row in _SCHED_VAL_ROWS]
pt = 0
for chunk_idx, size in enumerate(chunk_sizes):
output = SimpleNamespace(
input_token_logprobs=tuple(token_logprobs[pt : pt + size]),
input_top_logprobs_val=[_SCHED_VAL_ROWS[pt : pt + size]],
input_top_logprobs_idx=[_SCHED_IDX_ROWS[pt : pt + size]],
)
processor.add_input_logprob_return_values(
0,
req,
output,
0,
size,
last_prefill_chunk=chunk_idx == len(chunk_sizes) - 1,
)
pt += size
def test_flat_arrays_replace_nested_rows(self):
flag_off = self._make_req(flat=False)
self._run_prefill(flag_off, [3, 2])
flag_on = self._make_req(flat=True)
self._run_prefill(flag_on, [3, 2])
# Flag off: nested rows as today, no arrays.
self.assertIsNone(flag_off.logprob.input_top_logprobs_val_flat)
self.assertIsNone(flag_off.logprob.input_top_logprobs_flat_null_prefix)
self.assertEqual(
flag_off.logprob.input_top_logprobs_val, [None] + _SCHED_VAL_ROWS[:-1]
)
# Flag on: arrays carrying the nested rows' content, nested emptied.
val_arr = flag_on.logprob.input_top_logprobs_val_flat
idx_arr = flag_on.logprob.input_top_logprobs_idx_flat
self.assertEqual(val_arr.dtype, np.float32)
self.assertEqual(idx_arr.dtype, np.int32)
self.assertEqual(flag_on.logprob.input_top_logprobs_flat_null_prefix, 1)
np.testing.assert_array_equal(
val_arr, np.asarray(_SCHED_VAL_ROWS[:-1], dtype=np.float32)
)
np.testing.assert_array_equal(
idx_arr, np.asarray(_SCHED_IDX_ROWS[:-1], dtype=np.int32)
)
self.assertEqual(flag_on.logprob.input_top_logprobs_val, [])
self.assertEqual(flag_on.logprob.input_top_logprobs_idx, [])
# The non-top logprob results are untouched.
self.assertEqual(
flag_on.logprob.input_token_logprobs_val,
flag_off.logprob.input_token_logprobs_val,
)
self.assertEqual(
flag_on.logprob.input_token_logprobs_idx,
flag_off.logprob.input_token_logprobs_idx,
)
def test_chunked_matches_one_shot(self):
one_shot = self._make_req(flat=True)
self._run_prefill(one_shot, [5])
chunked = self._make_req(flat=True)
self._run_prefill(chunked, [2, 2, 1])
np.testing.assert_array_equal(
one_shot.logprob.input_top_logprobs_val_flat,
chunked.logprob.input_top_logprobs_val_flat,
)
np.testing.assert_array_equal(
one_shot.logprob.input_top_logprobs_idx_flat,
chunked.logprob.input_top_logprobs_idx_flat,
)
self.assertEqual(
one_shot.logprob.input_top_logprobs_flat_null_prefix,
chunked.logprob.input_top_logprobs_flat_null_prefix,
)
def test_unrepresentable_rows_fall_back_to_nested(self):
req = self._make_req(flat=True, num_tokens=3)
processor = _make_logprob_processor()
val_rows = [[-0.5, -2.5], [-0.25], [-0.125, -4.0]]
idx_rows = [[11, 22], [33], [55, 66]]
output = SimpleNamespace(
input_token_logprobs=(-0.5, -0.25, -0.125),
input_top_logprobs_val=[val_rows],
input_top_logprobs_idx=[idx_rows],
)
with self.assertLogs(
"sglang.srt.managers.scheduler_components.logprob_result_processor",
level="WARNING",
):
processor.add_input_logprob_return_values(
0, req, output, 0, 3, last_prefill_chunk=True
)
self.assertIsNone(req.logprob.input_top_logprobs_val_flat)
self.assertIsNone(req.logprob.input_top_logprobs_flat_null_prefix)
self.assertEqual(req.logprob.input_top_logprobs_val, [None] + val_rows[:-1])
self.assertEqual(req.logprob.input_top_logprobs_idx, [None] + idx_rows[:-1])
class TestFromArraysMatchesFromRows(CustomTestCase):
"""The tokenizer-manager from-arrays builder must produce the same
response fields as the rows-based builder."""
def _both(self, return_b64: bool):
from_rows = _build_flat_input_top_logprobs_fields(
_EXACT_VAL_ROWS, _IDX_ROWS, top_logprobs_num=2, return_b64=return_b64
)
val_arr, idx_arr, null_prefix = build_flat_input_top_logprobs_arrays(
_EXACT_VAL_ROWS, _IDX_ROWS, top_logprobs_num=2
)
from_arrays = _build_flat_input_top_logprobs_fields_from_arrays(
val_arr, idx_arr, null_prefix, return_b64=return_b64
)
return from_rows, from_arrays
def test_non_b64(self):
from_rows, from_arrays = self._both(return_b64=False)
self.assertEqual(from_rows, from_arrays)
def test_b64(self):
from_rows, from_arrays = self._both(return_b64=True)
self.assertEqual(from_rows, from_arrays)
def test_all_null_rows(self):
val_arr, idx_arr, null_prefix = build_flat_input_top_logprobs_arrays(
[None], [None], top_logprobs_num=2
)
self.assertEqual(val_arr.shape, (0, 2))
self.assertEqual(null_prefix, 1)
fields = _build_flat_input_top_logprobs_fields_from_arrays(
val_arr, idx_arr, null_prefix
)
self.assertEqual(
fields,
_build_flat_input_top_logprobs_fields([None], [None], top_logprobs_num=2),
)
class TestMetaInfoFromSchedulerArrays(CustomTestCase):
"""add_logprob_to_meta_info consumes scheduler-flat arrays directly."""
def _rows_meta(self, **state_kwargs) -> dict:
state = _make_state(
return_logprob=True,
top_logprobs_num=2,
return_flat_raw_top_logprobs=True,
**state_kwargs,
)
state.input_top_logprobs_val.extend(_EXACT_VAL_ROWS)
state.input_top_logprobs_idx.extend(_IDX_ROWS)
return _add_logprob_meta_info(state)
def _arrays_state(self, **state_kwargs) -> ReqState:
state = _make_state(
return_logprob=True,
top_logprobs_num=2,
return_flat_raw_top_logprobs=True,
**state_kwargs,
)
# Scheduler-flat requests arrive with empty nested rows and the arrays.
state.input_top_logprobs_scheduler_flat = build_flat_input_top_logprobs_arrays(
_EXACT_VAL_ROWS, _IDX_ROWS, top_logprobs_num=2
)
return state
def test_matches_rows_path_field_for_field(self):
got = _add_logprob_meta_info(self._arrays_state())
self.assertEqual(got, self._rows_meta())
def test_b64_matches_rows_path_field_for_field(self):
got = _add_logprob_meta_info(
self._arrays_state(return_flat_raw_top_logprobs_b64=True)
)
self.assertEqual(got, self._rows_meta(return_flat_raw_top_logprobs_b64=True))
def test_fields_cached_across_chunks(self):
state = self._arrays_state()
first = _add_logprob_meta_info(state)
again = _add_logprob_meta_info(state)
self.assertIs(
again["input_top_logprobs_val_flat"], first["input_top_logprobs_val_flat"]
)
def _make_batch_token_id_output(**overrides) -> BatchTokenIDOutput:
"""A two-request BatchTokenIDOutput with the required fields stubbed."""
n = 2
fields = dict(
rids=["r0", "r1"],
finished_reasons=[None] * n,
decoded_texts=["", ""],
decode_ids=[array("q", [1]), array("q", [2])],
read_offsets=[0] * n,
output_ids=None,
skip_special_tokens=[True] * n,
spaces_between_special_tokens=[True] * n,
no_stop_trim=[False] * n,
prompt_tokens=[5] * n,
reasoning_tokens=[0] * n,
completion_tokens=[1] * n,
cached_tokens=[0] * n,
input_token_logprobs_val=[[], []],
input_token_logprobs_idx=[[], []],
output_token_logprobs_val=[[], []],
output_token_logprobs_idx=[[], []],
input_top_logprobs_val=[[], []],
input_top_logprobs_idx=[[], []],
output_top_logprobs_val=[[], []],
output_top_logprobs_idx=[[], []],
input_token_ids_logprobs_val=[[], []],
input_token_ids_logprobs_idx=[[], []],
output_token_ids_logprobs_val=[[], []],
output_token_ids_logprobs_idx=[[], []],
output_token_entropy_val=None,
output_token_sampling_mask=None,
output_token_sampling_logprobs=None,
output_hidden_states=None,
routed_experts=None,
indexer_topk=None,
placeholder_tokens_idx=None,
placeholder_tokens_val=None,
)
fields.update(overrides)
return BatchTokenIDOutput(**fields)
class TestBatchOutputTransport(CustomTestCase):
"""The flat array fields must survive both IPC transports: pickle
(SGLANG_USE_PICKLE_IPC, the default) and msgpack (enc/dec hooks)."""
def _flat_output(self) -> BatchTokenIDOutput:
val_arr, idx_arr, null_prefix = build_flat_input_top_logprobs_arrays(
_EXACT_VAL_ROWS, _IDX_ROWS, top_logprobs_num=2
)
return _make_batch_token_id_output(
input_top_logprobs_val_flat=[None, val_arr],
input_top_logprobs_idx_flat=[None, idx_arr],
input_top_logprobs_flat_null_prefix=[None, null_prefix],
)
def _check_roundtrip(self, decoded, original):
self.assertIsNone(decoded.input_top_logprobs_val_flat[0])
self.assertIsNone(decoded.input_top_logprobs_idx_flat[0])
self.assertIsNone(decoded.input_top_logprobs_flat_null_prefix[0])
for got, sent in (
(
decoded.input_top_logprobs_val_flat[1],
original.input_top_logprobs_val_flat[1],
),
(
decoded.input_top_logprobs_idx_flat[1],
original.input_top_logprobs_idx_flat[1],
),
):
self.assertIsInstance(got, np.ndarray)
self.assertEqual(got.dtype, sent.dtype)
np.testing.assert_array_equal(got, sent)
self.assertEqual(decoded.input_top_logprobs_flat_null_prefix[1], 1)
def test_pickle_roundtrip(self):
output = self._flat_output()
decoded = pickle.loads(pickle.dumps(output, protocol=pickle.HIGHEST_PROTOCOL))
self._check_roundtrip(decoded, output)
def test_msgpack_roundtrip(self):
output = self._flat_output()
decoded = msgpack_decode(msgpack_encode(output))
self._check_roundtrip(decoded, output)
def test_fields_default_none(self):
output = _make_batch_token_id_output()
self.assertIsNone(output.input_top_logprobs_val_flat)
self.assertIsNone(output.input_top_logprobs_idx_flat)
self.assertIsNone(output.input_top_logprobs_flat_null_prefix)
@unittest.skipUnless( @unittest.skipUnless(
os.environ.get("SGLANG_BENCH_FLAT_RAW_TOP_LOGPROBS"), os.environ.get("SGLANG_BENCH_FLAT_RAW_TOP_LOGPROBS"),
"Serialization microbenchmark; set SGLANG_BENCH_FLAT_RAW_TOP_LOGPROBS=1 to run.", "Serialization microbenchmark; set SGLANG_BENCH_FLAT_RAW_TOP_LOGPROBS=1 to run.",
@@ -382,6 +713,41 @@ class BenchFlatRawTopLogprobsSerialization(CustomTestCase):
decode_b64, decode_b64,
) )
def test_bench_ipc_pickle(self):
"""Inter-process cost of BatchTokenIDOutput input-top fields: nested
per-position rows vs scheduler-flat arrays (two ZMQ pickle hops each
pay dumps + loads)."""
num_positions, k = 32768, 2
rng = np.random.default_rng(0)
vals = rng.standard_normal((num_positions, k)).astype(np.float32)
idxs = rng.integers(0, 150000, size=(num_positions, k), dtype=np.int32)
def best_of(fn, iters=10):
return min(
(lambda s=time.perf_counter(): (fn(), time.perf_counter() - s)[1])()
for _ in range(iters)
)
nested = _make_batch_token_id_output(
input_top_logprobs_val=[[None] + vals[1:].tolist(), []],
input_top_logprobs_idx=[[None] + idxs[1:].tolist(), []],
)
flat = _make_batch_token_id_output(
input_top_logprobs_val_flat=[vals[1:], None],
input_top_logprobs_idx_flat=[idxs[1:], None],
input_top_logprobs_flat_null_prefix=[1, None],
)
for name, obj in (("nested rows", nested), ("flat arrays", flat)):
payload = pickle.dumps(obj, protocol=pickle.HIGHEST_PROTOCOL)
dumps_ms = best_of(
lambda o=obj: pickle.dumps(o, protocol=pickle.HIGHEST_PROTOCOL)
)
loads_ms = best_of(lambda p=payload: pickle.loads(p))
print(
f"{name}: pickle.dumps {dumps_ms * 1e3:.2f} ms, "
f"pickle.loads {loads_ms * 1e3:.2f} ms, {len(payload) / 1e6:.2f} MB"
)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main(verbosity=2) unittest.main(verbosity=2)