Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 9 additions & 17 deletions python/sglang/srt/disaggregation/decode.py
Original file line number Diff line number Diff line change
Expand Up @@ -2554,23 +2554,15 @@ def _commit_transfer_to_req(self, decode_req: DecodeRequest):
"sampling mask buffer disabled on decode side"
)
sampling_mask_len = int(output_token_sampling_mask_len[0].item())
if sampling_mask_len < 0:
decode_req.req.output_token_sampling_mask.append(None)
decode_req.req.output_token_sampling_logprobs.append(None)
else:
decode_req.req.output_token_sampling_mask.append(
output_token_sampling_mask_idx[:sampling_mask_len].cpu().tolist()
)
if decode_req.req.sampling_logprobs_mode == "support":
decode_req.req.output_token_sampling_logprobs.append(
output_token_sampling_logprobs[:sampling_mask_len]
.cpu()
.tolist()
)
else:
decode_req.req.output_token_sampling_logprobs.append(
float(output_token_sampling_logprobs[0].item())
)
num_logprobs = (
sampling_mask_len
if decode_req.req.sampling_logprobs_mode == "support"
else 1
)
decode_req.req.sampling_mask_rows.append(
output_token_sampling_mask_idx[:sampling_mask_len].cpu().numpy(),
output_token_sampling_logprobs[:num_logprobs].cpu().numpy(),
)

decode_req.kv_receiver.clear()
decode_req.kv_receiver = None
Expand Down
56 changes: 11 additions & 45 deletions python/sglang/srt/disaggregation/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -518,51 +518,17 @@ def set_buf(self, req: Req):
device="cpu",
)
if req.return_sampling_mask:
# Sentinel -1: the decode side records None for this handoff token.
self.output_token_sampling_mask_len[req.metadata_buffer_index][0] = -1
sampling_masks = req.output_token_sampling_mask
sampling_logprobs = req.output_token_sampling_logprobs
if sampling_masks:
sampling_mask = sampling_masks[0]
sampling_logprobs_row = (
sampling_logprobs[0] if sampling_logprobs else None
)
if sampling_mask is not None and sampling_logprobs_row is not None:
mask_len = len(sampling_mask)
max_mask_len = self.output_token_sampling_mask_idx.shape[1]
if mask_len > max_mask_len:
raise RuntimeError(
f"Sampling mask length {mask_len} exceeds disaggregation "
f"metadata capacity {max_mask_len}. Increase "
"--sampling-mask-max-tokens."
)
self.output_token_sampling_mask_len[req.metadata_buffer_index][
0
] = mask_len
if mask_len:
self.output_token_sampling_mask_idx[
req.metadata_buffer_index, :mask_len
].copy_(
torch.tensor(
sampling_mask,
dtype=torch.int32,
device=self.output_token_sampling_mask_idx.device,
)
)
if req.sampling_logprobs_mode == "support" and mask_len:
self.output_token_sampling_logprobs[
req.metadata_buffer_index, :mask_len
].copy_(
torch.tensor(
sampling_logprobs_row,
dtype=torch.float32,
device=self.output_token_sampling_logprobs.device,
)
)
elif req.sampling_logprobs_mode == "selected":
self.output_token_sampling_logprobs[
req.metadata_buffer_index, 0
] = float(sampling_logprobs_row)
# Prefill streams a request only once its KV transfer ends or it aborts,
# so the first token's row is the only one queued here.
chunk = req.sampling_mask_rows.view()
mask_len = len(chunk.token_ids)
self.output_token_sampling_mask_len[req.metadata_buffer_index][0] = mask_len
self.output_token_sampling_mask_idx[
req.metadata_buffer_index, :mask_len
].copy_(torch.from_numpy(chunk.token_ids))
self.output_token_sampling_logprobs[
req.metadata_buffer_index, : len(chunk.logprobs)
].copy_(torch.from_numpy(chunk.logprobs))
# For PD + spec decode
if req.hidden_states_tensor is not None:
# speculative_eagle_topk should not be greater than 16 currently
Expand Down
7 changes: 3 additions & 4 deletions python/sglang/srt/layers/logits_processor.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
from enum import IntEnum
from typing import Any, Dict, List, Optional, Tuple, Union

import numpy as np
import torch
from torch import nn

Expand Down Expand Up @@ -229,10 +230,8 @@ class LogitsProcessorOutput:
# Post-filter support IDs and requested behavior logprobs, bounded by server
# capacity. Logprobs are normalized over the full realized support.
sampling_mask_output: Optional[SamplingMaskOutput] = None
next_token_sampling_mask_idx: Optional[List[Optional[List[int]]]] = None
next_token_sampling_logprobs: Optional[
List[Optional[Union[float, List[float]]]]
] = None
next_token_sampling_mask_idx: Optional[List[Optional[np.ndarray]]] = None
next_token_sampling_logprobs: Optional[List[Optional[np.ndarray]]] = None
next_token_sampling_mask_status: Optional[List[Optional[int]]] = None

## Part 3: Prefill-only. This part will be assigned in python/sglang/srt/layers/logits_processor.py::LogitsProcessor
Expand Down
1 change: 0 additions & 1 deletion python/sglang/srt/managers/detokenizer_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -495,7 +495,6 @@ def handle_batch_token_id_out(self, recv_obj: BatchTokenIDOutput):
output_token_ids_logprobs_idx=recv_obj.output_token_ids_logprobs_idx,
output_token_entropy_val=recv_obj.output_token_entropy_val,
output_token_sampling_mask=recv_obj.output_token_sampling_mask,
output_token_sampling_logprobs=recv_obj.output_token_sampling_logprobs,
output_hidden_states=recv_obj.output_hidden_states,
routed_experts=routed_experts,
indexer_topk=indexer_topk,
Expand Down
13 changes: 3 additions & 10 deletions python/sglang/srt/managers/io_struct.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,7 @@
get_return_hidden_states_mode,
)
from sglang.srt.multimodal.mm_utils import has_valid_data
from sglang.srt.sampling.sampling_mask import SamplingMaskChunk
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.utils import ImageData, VideoData
from sglang.srt.utils.field_validators import validate_optional_list_i64_1d_2d
Expand Down Expand Up @@ -1543,12 +1544,7 @@ class BatchTokenIDOutput(BaseBatchReq, kw_only=True):
output_token_ids_logprobs_val: TokenIdsLogprobValues
output_token_ids_logprobs_idx: TokenIdsLogprobIndices
output_token_entropy_val: Optional[List[Optional[float]]]
# Per-request chunks of output-token sampling supports. None when no request
# in the batch asks for return_sampling_mask.
output_token_sampling_mask: Optional[List[List[List[int]]]]
# Per-request chunks. Each output-token entry is a selected-token scalar or
# a list aligned with output_token_sampling_mask, according to the request.
output_token_sampling_logprobs: Optional[List[List[Union[float, List[float]]]]]
output_token_sampling_mask: Optional[List[Optional[SamplingMaskChunk]]]

# Hidden states
output_hidden_states: OutputHiddenStates
Expand Down Expand Up @@ -1643,10 +1639,7 @@ class BatchStrOutput(BaseBatchReq, kw_only=True):
output_token_ids_logprobs_val: TokenIdsLogprobValues
output_token_ids_logprobs_idx: TokenIdsLogprobIndices
output_token_entropy_val: Optional[List[Optional[float]]]
# Detokenizer pass-through for BatchTokenIDOutput.output_token_sampling_*;
# support-mode logprobs are aligned elementwise with the token IDs.
output_token_sampling_mask: Optional[List[List[List[int]]]]
output_token_sampling_logprobs: Optional[List[List[Union[float, List[float]]]]]
output_token_sampling_mask: Optional[List[Optional[SamplingMaskChunk]]]

# Hidden states
output_hidden_states: OutputHiddenStates
Expand Down
6 changes: 0 additions & 6 deletions python/sglang/srt/managers/multi_tokenizer_mixin.py
Original file line number Diff line number Diff line change
Expand Up @@ -247,9 +247,6 @@ def _handle_output_by_index(output, i):
output_token_sampling_mask=_extract_field_by_index(
output, "output_token_sampling_mask", i, check_length=False
),
output_token_sampling_logprobs=_extract_field_by_index(
output, "output_token_sampling_logprobs", i, check_length=False
),
output_hidden_states=_extract_field_by_index(
output, "output_hidden_states", i, check_length=False
),
Expand Down Expand Up @@ -368,9 +365,6 @@ def _handle_output_by_index(output, i):
output_token_sampling_mask=_extract_field_by_index(
output, "output_token_sampling_mask", i, check_length=False
),
output_token_sampling_logprobs=_extract_field_by_index(
output, "output_token_sampling_logprobs", i, check_length=False
),
output_hidden_states=_extract_field_by_index(
output, "output_hidden_states", i, check_length=False
),
Expand Down
11 changes: 4 additions & 7 deletions python/sglang/srt/managers/schedule_batch.py
Original file line number Diff line number Diff line change
Expand Up @@ -141,6 +141,7 @@
SchedulerReqTimeStats,
)
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.sampling.sampling_mask import SamplingMaskRows
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.utils import flatten_nested_list
from sglang.srt.utils.token_sequence_matcher import TokenSequenceMatcher
Expand Down Expand Up @@ -1220,7 +1221,6 @@ def __init__(
# TODO (Byron): send_output_token_logprobs_offset and send_decode_id_offset can be different in disaggregation mode
# because the decode server does not have the first output token logprobs
self.send_output_token_logprobs_offset: int = 0
self.send_output_sampling_mask_offset: int = 0

# Logprobs (arguments)
self.return_logprob = return_logprob
Expand Down Expand Up @@ -1256,12 +1256,9 @@ def __init__(
# Can contain either lists or GPU tensors (delayed copy optimization for prefill-only scoring)
self.logprob.output_token_ids_logprobs_val = []
self.logprob.output_token_ids_logprobs_idx = []
if return_sampling_mask:
self.output_token_sampling_mask = []
self.output_token_sampling_logprobs = []
else:
self.output_token_sampling_mask = None
self.output_token_sampling_logprobs = None
self.sampling_mask_rows: Optional[SamplingMaskRows] = (
SamplingMaskRows() if return_sampling_mask else None
)
self.hidden_states: List[List[float]] = []
self.hidden_states_tensor = None # Note: use tensor instead of list to transfer hidden_states when PD + MTP
self.output_topk_p = None
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1163,58 +1163,51 @@ def add_sampling_mask_return_values(
output: LogitsProcessorOutput,
) -> None:
"""Attach sparse sampling support metadata to the return values."""
mask = output.next_token_sampling_mask_idx
logprobs = output.next_token_sampling_logprobs
req.output_token_sampling_mask.append(None if mask is None else mask[i])
req.output_token_sampling_logprobs.append(
None if logprobs is None else logprobs[i]
req.sampling_mask_rows.append(
output.next_token_sampling_mask_idx[i],
output.next_token_sampling_logprobs[i],
)

@staticmethod
def materialize_sampling_mask_output(
reqs: List[Req],
output: Optional[LogitsProcessorOutput],
) -> None:
"""Convert opted-in tensor rows to batch-aligned Python results."""
"""Convert opted-in tensor rows to batch-aligned host rows."""
if output is None or output.sampling_mask_output is None:
return

sampling_output = output.sampling_mask_output
batch_indices = [i for i, req in enumerate(reqs) if req.return_sampling_mask]
lengths = sampling_output.lengths.tolist()
selected_logprobs = sampling_output.selected_logprobs.tolist()
statuses = sampling_output.statuses.tolist()
assert len(batch_indices) == len(lengths)

batch_size = len(reqs)
masks = [None] * batch_size
logprobs = [None] * batch_size
status_by_batch = [None] * batch_size
token_ids = sampling_output.token_ids.cpu()
token_ids = sampling_output.token_ids.cpu().numpy()
selected_logprobs = sampling_output.selected_logprobs.cpu().numpy()
support_logprobs = (
None
if sampling_output.support_logprobs is None
else sampling_output.support_logprobs.cpu()
else sampling_output.support_logprobs.cpu().numpy()
)
packed_width = token_ids.shape[1]
support_row = 0
for row, batch_index in enumerate(batch_indices):
returns_support_logprobs = (
reqs[batch_index].sampling_logprobs_mode == "support"
)
status = int(statuses[row])
length = int(lengths[row])
if status == SamplingMaskStatus.OK and not (0 <= length <= packed_width):
status = SamplingMaskStatus.INVALID
status_by_batch[batch_index] = status
if status == SamplingMaskStatus.OK:
masks[batch_index] = token_ids[row, :length].tolist()
masks[batch_index] = token_ids[row, :length]
if returns_support_logprobs:
logprobs[batch_index] = support_logprobs[
support_row, :length
].tolist()
logprobs[batch_index] = support_logprobs[support_row, :length]
else:
logprobs[batch_index] = float(selected_logprobs[row])
logprobs[batch_index] = selected_logprobs[row : row + 1]
if returns_support_logprobs:
support_row += 1

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
)
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
from sglang.srt.runtime_context import get_observability, get_parallel, get_serving
from sglang.srt.sampling.sampling_mask import SamplingMaskChunk
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils.weight_versions import compute_weight_version_spans
Expand Down Expand Up @@ -374,8 +375,7 @@ class _GenerationStreamAccumulator:
input_token_ids_logprobs_idx: Optional[list] = None
output_token_ids_logprobs_val: Optional[list] = None
output_token_ids_logprobs_idx: Optional[list] = None
output_token_sampling_mask: Optional[list[list[list[int]]]] = None
output_token_sampling_logprobs: Optional[list[list[float | list[float]]]] = None
output_token_sampling_mask: Optional[list[Optional[SamplingMaskChunk]]] = None
# Rust server mode: the Rust detokenizer reconstructs text/ids from the raw
# output tokens itself and never consumes the scheduler's incremental-detok
# offsets (decode_ids / read_offset), so that per-step bookkeeping is skipped.
Expand Down Expand Up @@ -407,7 +407,6 @@ def __post_init__(self) -> None:
self.output_token_ids_logprobs_idx = []
if self.return_sampling_mask:
self.output_token_sampling_mask = []
self.output_token_sampling_logprobs = []

def _beam_admits(self, *, req: Req) -> bool:
# Only the leader is ever streamed, and only at group finish.
Expand Down Expand Up @@ -624,22 +623,9 @@ def accept(self, *, req: Req) -> None:

if self.return_sampling_mask:
if req.return_sampling_mask:
send_output_sampling_mask_offset = req.send_output_sampling_mask_offset
sampling_mask_end = len(req.output_token_sampling_mask)
self.output_token_sampling_mask.append(
req.output_token_sampling_mask[
send_output_sampling_mask_offset:sampling_mask_end
]
)
self.output_token_sampling_logprobs.append(
req.output_token_sampling_logprobs[
send_output_sampling_mask_offset:sampling_mask_end
]
)
req.send_output_sampling_mask_offset = sampling_mask_end
self.output_token_sampling_mask.append(req.sampling_mask_rows.take())
else:
self.output_token_sampling_mask.append([])
self.output_token_sampling_logprobs.append([])
self.output_token_sampling_mask.append(None)

if self.return_hidden_states:
if req.return_hidden_states:
Expand Down Expand Up @@ -748,7 +734,6 @@ def to_payload(
output_token_ids_logprobs_idx=self.output_token_ids_logprobs_idx,
output_token_entropy_val=None,
output_token_sampling_mask=self.output_token_sampling_mask,
output_token_sampling_logprobs=self.output_token_sampling_logprobs,
output_hidden_states=self.output_hidden_states,
routed_experts=self.routed_experts,
indexer_topk=self.indexer_topk,
Expand Down
11 changes: 5 additions & 6 deletions python/sglang/srt/managers/tokenizer_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -2397,12 +2397,11 @@ async def _handle_batch_output(
):
output_sampling_mask = recv_obj.output_token_sampling_mask
if output_sampling_mask is not None:
state.output_token_sampling_mask.extend(output_sampling_mask[i])
output_sampling_logprobs = recv_obj.output_token_sampling_logprobs
if output_sampling_logprobs is not None:
state.output_token_sampling_logprobs.extend(
output_sampling_logprobs[i]
)
masks, logprobs = output_sampling_mask[i].to_lists(
support_logprobs=state.obj.sampling_logprobs_mode == "support"
)
state.output_token_sampling_mask.extend(masks)
state.output_token_sampling_logprobs.extend(logprobs)
meta_info["output_token_sampling_mask"] = (
state.output_token_sampling_mask
)
Expand Down
Loading
Loading