[perf] Assemble flat prompt top logprobs scheduler-side as numpy arrays (#32223)
This commit is contained in:
@@ -462,6 +462,9 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
||||
output_token_logprobs_idx=recv_obj.output_token_logprobs_idx,
|
||||
input_top_logprobs_val=recv_obj.input_top_logprobs_val,
|
||||
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_idx=recv_obj.output_top_logprobs_idx,
|
||||
input_token_ids_logprobs_val=recv_obj.input_token_ids_logprobs_val,
|
||||
|
||||
@@ -38,6 +38,7 @@ from typing import (
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
Tuple,
|
||||
Type,
|
||||
Union,
|
||||
)
|
||||
@@ -831,6 +832,10 @@ class TokenizedGenerateReqInput(BaseReq, kw_only=True):
|
||||
stream: bool
|
||||
# Whether to return sparse output-token support from top-k/top-p/min-p sampling.
|
||||
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
|
||||
return_hidden_states: bool = False
|
||||
@@ -1229,6 +1234,39 @@ CachedTokensDetails = Dict[str, Union[int, str]]
|
||||
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):
|
||||
# The finish reason
|
||||
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_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):
|
||||
# The finish reason
|
||||
@@ -1401,6 +1448,12 @@ class BatchStrOutput(BaseBatchReq, kw_only=True):
|
||||
spec_correct_drafts_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):
|
||||
# The finish reason
|
||||
|
||||
@@ -211,6 +211,15 @@ def _handle_output_by_index(output, i):
|
||||
input_top_logprobs_idx=_extract_field_by_index(
|
||||
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, "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(
|
||||
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, "output_top_logprobs_val", i, check_length=False
|
||||
),
|
||||
|
||||
@@ -687,6 +687,12 @@ class ReqLogprob:
|
||||
input_token_logprobs_idx: Optional[List[int]] = None
|
||||
input_top_logprobs_val: Optional[List[List[float]]] = 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_idx: Optional[List[List[int]]] = None
|
||||
output_token_logprobs_val: Optional[list] = None
|
||||
@@ -725,6 +731,7 @@ class Req(ReqDllmMixin):
|
||||
dllm_config: Optional[DllmConfig] = None,
|
||||
token_ids_logprob: List[int] = None,
|
||||
return_sampling_mask: bool = False,
|
||||
return_flat_raw_top_logprobs: bool = False,
|
||||
stream: bool = False,
|
||||
origin_input_ids_unpadded: Optional[array[int]] = None,
|
||||
lora_id: Optional[str] = None,
|
||||
@@ -943,6 +950,7 @@ class Req(ReqDllmMixin):
|
||||
self.temp_scaled_logprobs = False
|
||||
self.top_p_normalized_logprobs = False
|
||||
self.return_sampling_mask = return_sampling_mask
|
||||
self.return_flat_raw_top_logprobs = return_flat_raw_top_logprobs
|
||||
|
||||
# Logprobs (return values)
|
||||
# True means the input logprob has been already sent to detokenizer.
|
||||
|
||||
@@ -2242,6 +2242,7 @@ class Scheduler(
|
||||
top_logprobs_num=recv_req.top_logprobs_num,
|
||||
token_ids_logprob=recv_req.token_ids_logprob,
|
||||
return_sampling_mask=recv_req.return_sampling_mask,
|
||||
return_flat_raw_top_logprobs=recv_req.return_flat_raw_top_logprobs,
|
||||
stream=recv_req.stream,
|
||||
lora_id=recv_req.lora_id,
|
||||
session_id=recv_req.session_id,
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import (
|
||||
List,
|
||||
@@ -10,6 +11,7 @@ import torch
|
||||
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
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.runtime_context import get_exec
|
||||
from sglang.srt.server_args import (
|
||||
@@ -17,6 +19,8 @@ from sglang.srt.server_args import (
|
||||
ServerArgs,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(kw_only=True, slots=True, frozen=True)
|
||||
class SchedulerLogprobResultProcessor:
|
||||
@@ -84,6 +88,35 @@ class SchedulerLogprobResultProcessor:
|
||||
req.temp_input_top_logprobs_idx = 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:
|
||||
"""Process input token IDs logprobs."""
|
||||
if req.logprob.token_ids_logprob is None:
|
||||
@@ -265,6 +298,10 @@ class SchedulerLogprobResultProcessor:
|
||||
== 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(
|
||||
self,
|
||||
i: int,
|
||||
|
||||
@@ -312,6 +312,12 @@ class _GenerationStreamAccumulator:
|
||||
output_token_logprobs_idx: Optional[list] = None
|
||||
input_top_logprobs_val: 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_idx: Optional[list] = None
|
||||
input_token_ids_logprobs_val: Optional[list] = None
|
||||
@@ -340,6 +346,9 @@ class _GenerationStreamAccumulator:
|
||||
self.output_token_logprobs_idx = []
|
||||
self.input_top_logprobs_val = []
|
||||
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_idx = []
|
||||
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_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(
|
||||
req.logprob.input_token_ids_logprobs_val
|
||||
)
|
||||
@@ -476,6 +496,9 @@ class _GenerationStreamAccumulator:
|
||||
self.input_token_logprobs_idx.append([])
|
||||
self.input_top_logprobs_val.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_idx.append([])
|
||||
|
||||
@@ -613,6 +636,23 @@ class _GenerationStreamAccumulator:
|
||||
output_token_logprobs_idx=self.output_token_logprobs_idx,
|
||||
input_top_logprobs_val=self.input_top_logprobs_val,
|
||||
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_idx=self.output_top_logprobs_idx,
|
||||
input_token_ids_logprobs_val=self.input_token_ids_logprobs_val,
|
||||
|
||||
@@ -84,6 +84,7 @@ from sglang.srt.managers.io_struct import (
|
||||
UpdateWeightFromDiskReqOutput,
|
||||
async_sock_recv,
|
||||
async_sock_send,
|
||||
build_flat_input_top_logprobs_arrays,
|
||||
sock_send,
|
||||
unwrap_from_pickle,
|
||||
)
|
||||
@@ -233,6 +234,12 @@ class ReqState:
|
||||
# prefill chunks arrive, so streaming decode chunks reuse the payload.
|
||||
input_top_logprobs_flat_fields: Optional[Dict[str, Any]] = None
|
||||
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
|
||||
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
|
||||
widths change later without a wire break.
|
||||
"""
|
||||
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:
|
||||
# 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})."
|
||||
)
|
||||
val_arr, idx_arr, null_prefix = build_flat_input_top_logprobs_arrays(
|
||||
input_top_logprobs_val, input_top_logprobs_idx, top_logprobs_num
|
||||
)
|
||||
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] = {}
|
||||
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(
|
||||
val_arr.tobytes()
|
||||
).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_idx_flat_b64_dtype"] = "int32"
|
||||
else:
|
||||
fields["input_top_logprobs_val_flat"] = [v for row in val_rows for v in row]
|
||||
fields["input_top_logprobs_idx_flat"] = [i for row in idx_rows for i in row]
|
||||
fields["input_top_logprobs_shape"] = [len(val_rows), k]
|
||||
fields["input_top_logprobs_val_flat"] = val_arr.reshape(-1).tolist()
|
||||
fields["input_top_logprobs_idx_flat"] = idx_arr.reshape(-1).tolist()
|
||||
fields["input_top_logprobs_shape"] = [val_arr.shape[0], val_arr.shape[1]]
|
||||
fields["input_top_logprobs_null_prefix"] = null_prefix
|
||||
return fields
|
||||
|
||||
@@ -1264,6 +1283,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
top_logprobs_num=obj.top_logprobs_num,
|
||||
token_ids_logprob=obj.token_ids_logprob,
|
||||
return_sampling_mask=obj.return_sampling_mask,
|
||||
return_flat_raw_top_logprobs=obj.return_flat_raw_top_logprobs,
|
||||
stream=obj.stream,
|
||||
rid=obj.rid,
|
||||
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
|
||||
# GenerateReqInput here.
|
||||
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.
|
||||
if state.input_top_logprobs_flat_num_rows != len(
|
||||
state.input_top_logprobs_val
|
||||
@@ -2424,6 +2460,15 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
state.input_top_logprobs_idx.extend(
|
||||
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(
|
||||
recv_obj.output_top_logprobs_val[recv_obj_index]
|
||||
)
|
||||
|
||||
@@ -6,8 +6,11 @@ import asyncio
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
import pickle
|
||||
import time
|
||||
import unittest
|
||||
from array import array
|
||||
from types import SimpleNamespace
|
||||
|
||||
import numpy as np
|
||||
|
||||
@@ -16,13 +19,25 @@ from sglang.test.test_utils import CustomTestCase, 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 (
|
||||
ReqState,
|
||||
TokenizerManager,
|
||||
_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.sampling.sampling_params import SamplingParams
|
||||
|
||||
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]]
|
||||
_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:
|
||||
"""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(
|
||||
os.environ.get("SGLANG_BENCH_FLAT_RAW_TOP_LOGPROBS"),
|
||||
"Serialization microbenchmark; set SGLANG_BENCH_FLAT_RAW_TOP_LOGPROBS=1 to run.",
|
||||
@@ -382,6 +713,41 @@ class BenchFlatRawTopLogprobsSerialization(CustomTestCase):
|
||||
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__":
|
||||
unittest.main(verbosity=2)
|
||||
|
||||
Reference in New Issue
Block a user