From 58974ca16ca2a4bb2f02f9ceb9622a0fd2ccf7f8 Mon Sep 17 00:00:00 2001 From: Sam Shleifer Date: Fri, 31 Jul 2026 21:00:51 -0400 Subject: [PATCH] [perf] Assemble flat prompt top logprobs scheduler-side as numpy arrays (#32223) --- .../srt/managers/detokenizer_manager.py | 3 + python/sglang/srt/managers/io_struct.py | 53 +++ .../srt/managers/multi_tokenizer_mixin.py | 18 + python/sglang/srt/managers/schedule_batch.py | 8 + python/sglang/srt/managers/scheduler.py | 1 + .../logprob_result_processor.py | 37 ++ .../scheduler_components/output_streamer.py | 40 ++ .../sglang/srt/managers/tokenizer_manager.py | 87 ++++- .../managers/test_flat_raw_top_logprobs.py | 368 +++++++++++++++++- 9 files changed, 593 insertions(+), 22 deletions(-) diff --git a/python/sglang/srt/managers/detokenizer_manager.py b/python/sglang/srt/managers/detokenizer_manager.py index 7b3ca52e5..19a7f4938 100644 --- a/python/sglang/srt/managers/detokenizer_manager.py +++ b/python/sglang/srt/managers/detokenizer_manager.py @@ -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, diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 32486141c..06282f331 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -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 diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py index dcd88eb13..0ea3e5065 100644 --- a/python/sglang/srt/managers/multi_tokenizer_mixin.py +++ b/python/sglang/srt/managers/multi_tokenizer_mixin.py @@ -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 ), diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index d63fb2af6..6f37e9c3e 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -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. diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 8857ae8cb..50fd2e9e5 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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, diff --git a/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py b/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py index 8c72deffd..f1f659565 100644 --- a/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py @@ -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, diff --git a/python/sglang/srt/managers/scheduler_components/output_streamer.py b/python/sglang/srt/managers/scheduler_components/output_streamer.py index 1a61e7555..9f5b2329a 100644 --- a/python/sglang/srt/managers/scheduler_components/output_streamer.py +++ b/python/sglang/srt/managers/scheduler_components/output_streamer.py @@ -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, diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 0d7237db8..bd7fc94ee 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -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] ) diff --git a/test/registered/unit/managers/test_flat_raw_top_logprobs.py b/test/registered/unit/managers/test_flat_raw_top_logprobs.py index 3dab14a17..d4225dd9a 100644 --- a/test/registered/unit/managers/test_flat_raw_top_logprobs.py +++ b/test/registered/unit/managers/test_flat_raw_top_logprobs.py @@ -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)