Return top-p/top-k sampling mask/nucleas (#27408)

This commit is contained in:
hari-hm
2026-07-13 17:58:51 -05:00
committed by GitHub
parent 0ee236ebdf
commit a909077d22
20 changed files with 719 additions and 36 deletions
@@ -1634,6 +1634,9 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
output_token_logprobs_idx,
output_top_logprobs_val,
output_top_logprobs_idx,
output_token_sampling_mask_len,
output_token_sampling_mask_idx,
output_token_sampling_logprobs,
output_topk_p,
output_topk_index,
output_hidden_states,
@@ -1748,6 +1751,21 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
: decode_req.req.logprob.top_logprobs_num
].tolist()
)
if decode_req.req.return_sampling_mask:
assert (
output_token_sampling_mask_idx is not None
), "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()
)
decode_req.req.output_token_sampling_logprobs.append(
float(output_token_sampling_logprobs[0].item())
)
decode_req.kv_receiver.clear()
decode_req.kv_receiver = None
@@ -674,6 +674,10 @@ class SchedulerDisaggregationPrefillMixin:
logits_output,
)
logprob_pt += num_input_logprobs
if req.return_sampling_mask:
self.batch_result_processor.add_sampling_mask_return_values(
i, req, logits_output
)
if not req.pending_bootstrap:
self.send_kv_chunk(req, last_chunk=True)
req.time_stats.set_prefill_transfer_queue_entry_time()
+101 -35
View File
@@ -226,9 +226,15 @@ class MetadataBuffers:
hidden_size: int,
hidden_states_dtype: torch.dtype,
max_top_logprobs_num: int = 128,
max_sampling_mask_tokens: Optional[int] = None,
custom_mem_pool: torch.cuda.MemPool = None,
):
self.custom_mem_pool = custom_mem_pool
if max_sampling_mask_tokens is None:
max_sampling_mask_tokens = (
envs.SGLANG_DISAGGREGATION_SAMPLING_MASK_MAX_TOKENS.get()
)
self.enable_sampling_mask = max_sampling_mask_tokens > 0
bootstrap_room_dtype = torch.uint64
device = "cpu"
if is_npu():
@@ -266,6 +272,19 @@ class MetadataBuffers:
self.output_top_logprobs_idx = torch.zeros(
(size, max_top_logprobs_num), dtype=torch.int32, device=device
)
self.output_token_sampling_mask_len = None
self.output_token_sampling_mask_idx = None
self.output_token_sampling_logprobs = None
if self.enable_sampling_mask:
self.output_token_sampling_mask_len = torch.zeros(
(size, 16), dtype=torch.int32, device=device
)
self.output_token_sampling_mask_idx = torch.zeros(
(size, max_sampling_mask_tokens), dtype=torch.int32, device=device
)
self.output_token_sampling_logprobs = torch.zeros(
(size, 16), dtype=torch.float32, device=device
)
# For PD + spec decode
self.output_topk_p = torch.zeros(
(size, 16), dtype=torch.float32, device=device
@@ -282,45 +301,43 @@ class MetadataBuffers:
)
def get_buf_infos(self):
ptrs = [
self.output_ids.data_ptr(),
self.cached_tokens.data_ptr(),
self.output_token_logprobs_val.data_ptr(),
self.output_token_logprobs_idx.data_ptr(),
self.output_top_logprobs_val.data_ptr(),
self.output_top_logprobs_idx.data_ptr(),
self.output_topk_p.data_ptr(),
self.output_topk_index.data_ptr(),
self.output_hidden_states.data_ptr(),
self.bootstrap_room.data_ptr(),
]
data_lens = [
self.output_ids.nbytes,
self.cached_tokens.nbytes,
self.output_token_logprobs_val.nbytes,
self.output_token_logprobs_idx.nbytes,
self.output_top_logprobs_val.nbytes,
self.output_top_logprobs_idx.nbytes,
self.output_topk_p.nbytes,
self.output_topk_index.nbytes,
self.output_hidden_states.nbytes,
self.bootstrap_room.nbytes,
]
item_lens = [
self.output_ids[0].nbytes,
self.cached_tokens[0].nbytes,
self.output_token_logprobs_val[0].nbytes,
self.output_token_logprobs_idx[0].nbytes,
self.output_top_logprobs_val[0].nbytes,
self.output_top_logprobs_idx[0].nbytes,
self.output_topk_p[0].nbytes,
self.output_topk_index[0].nbytes,
self.output_hidden_states[0].nbytes,
self.bootstrap_room[0].nbytes,
bufs = [
self.output_ids,
self.cached_tokens,
self.output_token_logprobs_val,
self.output_token_logprobs_idx,
self.output_top_logprobs_val,
self.output_top_logprobs_idx,
]
if self.enable_sampling_mask:
bufs.extend(
[
self.output_token_sampling_mask_len,
self.output_token_sampling_mask_idx,
self.output_token_sampling_logprobs,
]
)
bufs.extend(
[
self.output_topk_p,
self.output_topk_index,
self.output_hidden_states,
self.bootstrap_room,
]
)
ptrs = [buf.data_ptr() for buf in bufs]
data_lens = [buf.nbytes for buf in bufs]
item_lens = [buf[0].nbytes for buf in bufs]
return ptrs, data_lens, item_lens
def get_buf(self, idx: int):
sampling_mask_len = None
sampling_mask_idx = None
sampling_logprobs = None
if self.enable_sampling_mask:
sampling_mask_len = self.output_token_sampling_mask_len[idx].clone()
sampling_mask_idx = self.output_token_sampling_mask_idx[idx].clone()
sampling_logprobs = self.output_token_sampling_logprobs[idx].clone()
return (
self.output_ids[idx].clone(),
self.cached_tokens[idx].clone(),
@@ -328,6 +345,9 @@ class MetadataBuffers:
self.output_token_logprobs_idx[idx].clone(),
self.output_top_logprobs_val[idx].clone(),
self.output_top_logprobs_idx[idx].clone(),
sampling_mask_len,
sampling_mask_idx,
sampling_logprobs,
self.output_topk_p[idx].clone(),
self.output_topk_index[idx].clone(),
self.output_hidden_states[idx].clone(),
@@ -366,6 +386,14 @@ class MetadataBuffers:
)
if req.logprob.output_top_logprobs_val: # not none or empty list
top_logprobs_len = len(req.logprob.output_top_logprobs_val[0])
max_top_logprobs_len = self.output_top_logprobs_val.shape[1]
if top_logprobs_len > max_top_logprobs_len:
raise RuntimeError(
f"top_logprobs_num {top_logprobs_len} exceeds "
f"disaggregation metadata capacity {max_top_logprobs_len}. "
"Lower top_logprobs_num or increase the metadata buffer."
)
self.output_top_logprobs_val[req.metadata_buffer_index][
: len(req.logprob.output_top_logprobs_val[0])
] = torch.tensor(
@@ -381,6 +409,44 @@ class MetadataBuffers:
dtype=torch.int32,
device="cpu",
)
if req.return_sampling_mask:
if not self.enable_sampling_mask:
raise RuntimeError(
"return_sampling_mask with disaggregation requires "
"SGLANG_DISAGGREGATION_SAMPLING_MASK_MAX_TOKENS > 0."
)
# 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_logprob = sampling_logprobs[0] if sampling_logprobs else None
if sampling_mask is not None and sampling_logprob 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 "
"SGLANG_DISAGGREGATION_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,
)
)
self.output_token_sampling_logprobs[req.metadata_buffer_index][
0
] = float(sampling_logprob)
# For PD + spec decode
if req.hidden_states_tensor is not None:
# speculative_eagle_topk should not be greater than 16 currently
+1
View File
@@ -352,6 +352,7 @@ class Envs:
SGLANG_DISAGG_PREFILL_EARLY_SEND_CACHED_PREFIX = EnvBool(True)
SGLANG_DISAGGREGATION_ALL_CP_RANKS_TRANSFER = EnvBool(False)
SGLANG_DISAGGREGATION_FORCE_QUERY_PREFILL_DP_RANK = EnvBool(False)
SGLANG_DISAGGREGATION_SAMPLING_MASK_MAX_TOKENS = EnvInt(0)
# Scheduler: others:
# in seconds. Set if you observe high memory accumulation over a long serving period.
@@ -174,6 +174,10 @@ class LogitsProcessorOutput:
List[Union[List[float], torch.Tensor]]
] = None
next_token_token_ids_logprobs_idx: Optional[List] = None
# Sparse top-k/top-p/min-p support ids and selected-token logprob after
# truncation/renormalization. Only populated when requested.
next_token_sampling_mask_idx: Optional[List[Optional[List[int]]]] = None
next_token_sampling_logprobs: Optional[List[Optional[float]]] = None
## Part 3: Prefill-only. This part will be assigned in python/sglang/srt/layers/logits_processor.py::LogitsProcessor
# The logprobs of input tokens. shape: [#token]
+129
View File
@@ -114,6 +114,7 @@ class Sampler(nn.Module):
# Preprocess logits (custom processors and NaN handling)
logits = self._preprocess_logits(logits, sampling_info)
return_sampling_mask = any(sampling_info.return_sampling_masks or [])
if sampling_info.is_all_greedy:
if _use_aiter and not _disable_aiter_greedy_sample:
@@ -123,6 +124,10 @@ class Sampler(nn.Module):
_aiter_greedy_sample(batch_next_token_ids, logits)
else:
batch_next_token_ids = torch.argmax(logits, -1)
if return_sampling_mask:
self._attach_greedy_sampling_mask_to_output(
logits_output, sampling_info, batch_next_token_ids
)
if return_logprob:
original_logprobs = logprobs = torch.nn.functional.log_softmax(
logits, dim=-1
@@ -183,6 +188,16 @@ class Sampler(nn.Module):
batch_next_token_ids = self._sample_from_probs(
probs, sampling_info, positions, simple_sampling_case
)
if return_sampling_mask:
sampling_mask_data = self._compute_sampling_mask_from_probs(
probs, sampling_info
)
self._attach_sampling_mask_to_output(
logits_output,
sampling_info,
batch_next_token_ids,
sampling_mask_data,
)
if return_logprob and not SGLANG_RETURN_ORIGINAL_LOGPROB:
logprobs = (
logprobs_via_logsoftmax_kernel
@@ -260,6 +275,120 @@ class Sampler(nn.Module):
raise ValueError(f"Invalid sampling backend: {backend}")
return batch_next_token_ids
def _compute_sampling_mask_from_probs(
self, probs: torch.Tensor, sampling_info: SamplingBatchInfo
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
"""Return sorted token ids, sorted probs, keep mask, and raw probs."""
vocab_size = probs.shape[-1]
max_top_k = sampling_info.sampling_mask_max_top_k
if 0 < max_top_k < vocab_size:
probs_sort, probs_idx = torch.topk(
probs,
k=max_top_k,
dim=-1,
largest=True,
sorted=True,
)
positions = torch.arange(max_top_k, device=probs.device).view(1, -1)
else:
probs_sort, probs_idx = probs.sort(dim=-1, descending=True)
positions = torch.arange(vocab_size, device=probs.device).view(1, -1)
probs_sum = torch.cumsum(probs_sort, dim=-1)
keep_mask = positions < sampling_info.top_ks.view(-1, 1)
keep_mask &= (probs_sum - probs_sort) <= sampling_info.top_ps.view(-1, 1)
if sampling_info.need_min_p_sampling:
min_p_thresholds = probs_sort[:, 0] * sampling_info.min_ps
keep_mask &= probs_sort >= min_p_thresholds.view(-1, 1)
return probs_idx, probs_sort, keep_mask, probs
def _attach_greedy_sampling_mask_to_output(
self,
logits_output: LogitsProcessorOutput,
sampling_info: SamplingBatchInfo,
batch_next_token_ids: torch.Tensor,
) -> None:
tokens = batch_next_token_ids.to(torch.int32).cpu().tolist()
masks = []
logprobs = []
for i, should_return in enumerate(sampling_info.return_sampling_masks or []):
if should_return:
masks.append([int(tokens[i])])
logprobs.append(0.0)
else:
masks.append(None)
logprobs.append(None)
logits_output.next_token_sampling_mask_idx = masks
logits_output.next_token_sampling_logprobs = logprobs
def _attach_sampling_mask_to_output(
self,
logits_output: LogitsProcessorOutput,
sampling_info: SamplingBatchInfo,
batch_next_token_ids: torch.Tensor,
sampling_mask_data: Tuple[
torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor
],
) -> None:
probs_idx, probs_sort, keep_mask, probs = sampling_mask_data
return_sampling_masks = sampling_info.return_sampling_masks or []
if not return_sampling_masks:
logits_output.next_token_sampling_mask_idx = []
logits_output.next_token_sampling_logprobs = []
return
sampled_tokens = batch_next_token_ids.view(-1, 1)
sampled_matches_all = probs_idx == sampled_tokens
sampled_in_idx = sampled_matches_all.any(dim=-1)
# The sampler is the source of truth for the rollout action space. If a
# backend/numeric edge chooses a token just outside the reconstructed
# prefix, include that sampled token so training can replay a support
# that contained the rollout action.
effective_keep_mask = keep_mask | sampled_matches_all
selected_raw_probs = torch.gather(probs, 1, sampled_tokens).squeeze(1)
support_mass = torch.where(
effective_keep_mask, probs_sort, torch.zeros_like(probs_sort)
).sum(dim=-1)
support_mass = support_mass + torch.where(
sampled_in_idx, torch.zeros_like(selected_raw_probs), selected_raw_probs
)
selected_logprobs = torch.log(
selected_raw_probs.float()
/ support_mass.float().clamp_min(torch.finfo(torch.float32).tiny)
)
flat_rows, flat_cols = effective_keep_mask.nonzero(as_tuple=True)
flat_ids = probs_idx[flat_rows, flat_cols].to(torch.int32)
mask_lengths = effective_keep_mask.sum(dim=-1, dtype=torch.int32)
flat_ids_cpu = flat_ids.cpu().tolist()
mask_lengths_cpu = mask_lengths.cpu().tolist()
sampled_in_idx_cpu = sampled_in_idx.cpu().tolist()
sampled_tokens_cpu = batch_next_token_ids.to(torch.int32).cpu().tolist()
selected_logprobs_cpu = selected_logprobs.cpu().tolist()
masks = []
logprobs = []
cursor = 0
for i, should_return in enumerate(return_sampling_masks):
mask_len = int(mask_lengths_cpu[i])
row_ids = flat_ids_cpu[cursor : cursor + mask_len]
cursor += mask_len
if not sampled_in_idx_cpu[i]:
row_ids.append(int(sampled_tokens_cpu[i]))
if should_return:
masks.append(row_ids)
logprobs.append(float(selected_logprobs_cpu[i]))
else:
masks.append(None)
logprobs.append(None)
logits_output.next_token_sampling_mask_idx = masks
logits_output.next_token_sampling_logprobs = logprobs
def _sample_from_logprobs(
self,
logprobs: torch.Tensor,
@@ -445,6 +445,8 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
output_token_ids_logprobs_val=recv_obj.output_token_ids_logprobs_val,
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,
+20
View File
@@ -204,6 +204,8 @@ class GenerateReqInput:
top_logprobs_num: Optional[Union[List[int], int]] = None
# If return logprobs, the token ids to return logprob for.
token_ids_logprob: Optional[Union[List[List[int]], List[int]]] = None
# Whether to return output-token sampling support and renormalized logprobs.
return_sampling_mask: Optional[Union[List[bool], bool]] = None
# Whether to detokenize tokens in text in the returned logprobs.
return_text_in_logprobs: bool = False
# Whether to stream output.
@@ -436,6 +438,8 @@ class GenerateReqInput:
self.top_logprobs_num = 0
if not self.token_ids_logprob: # covers both None and []
self.token_ids_logprob = None
if self.return_sampling_mask is None:
self.return_sampling_mask = False
def _normalize_batch_inputs(self):
"""Normalize inputs for a batch of examples, including parallel sampling expansion."""
@@ -600,6 +604,9 @@ class GenerateReqInput:
self.top_logprobs_num = normalize_param(
self.top_logprobs_num, 0, "top_logprobs_num"
)
self.return_sampling_mask = normalize_param(
self.return_sampling_mask, False, "return_sampling_mask"
)
# Handle token_ids_logprob specially due to its nested structure
if not self.token_ids_logprob: # covers both None and []
@@ -715,6 +722,7 @@ class GenerateReqInput:
logprob_start_len=self.logprob_start_len[i],
top_logprobs_num=self.top_logprobs_num[i],
token_ids_logprob=self.token_ids_logprob[i],
return_sampling_mask=self.return_sampling_mask[i],
return_text_in_logprobs=self.return_text_in_logprobs,
stream=self.stream,
log_metrics=self.log_metrics,
@@ -798,6 +806,8 @@ class TokenizedGenerateReqInput(BaseReq, kw_only=True):
token_ids_logprob: Optional[List[int]]
# Whether to stream output
stream: bool
# Whether to return sparse output-token support from top-k/top-p/min-p sampling.
return_sampling_mask: bool = False
# Whether to return hidden states
return_hidden_states: bool = False
@@ -1230,6 +1240,12 @@ 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]]
# Per-request chunks of selected-token logprobs renormalized over the
# corresponding sampling supports. None when sampling masks are not returned.
output_token_sampling_logprobs: Optional[List[List]]
# Hidden states
output_hidden_states: OutputHiddenStates
@@ -1309,6 +1325,10 @@ 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_*.
# None when sampling masks are not returned.
output_token_sampling_mask: Optional[List[List]]
output_token_sampling_logprobs: Optional[List[List]]
# Hidden states
output_hidden_states: OutputHiddenStates
@@ -232,6 +232,12 @@ def _handle_output_by_index(output, i):
output_token_entropy_val=_extract_field_by_index(
output, "output_token_entropy_val", i, check_length=False
),
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
),
@@ -334,6 +340,12 @@ def _handle_output_by_index(output, i):
output_token_entropy_val=_extract_field_by_index(
output, "output_token_entropy_val", i, check_length=False
),
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
),
@@ -678,6 +678,7 @@ class Req(ReqDllmMixin):
top_logprobs_num: int = 0,
dllm_config: Optional[DllmConfig] = None,
token_ids_logprob: List[int] = None,
return_sampling_mask: bool = False,
stream: bool = False,
origin_input_ids_unpadded: Optional[array[int]] = None,
lora_id: Optional[str] = None,
@@ -888,6 +889,7 @@ class Req(ReqDllmMixin):
# 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
@@ -897,6 +899,9 @@ class Req(ReqDllmMixin):
top_logprobs_num=top_logprobs_num,
token_ids_logprob=token_ids_logprob,
)
self.temp_scaled_logprobs = False
self.top_p_normalized_logprobs = False
self.return_sampling_mask = return_sampling_mask
# Logprobs (return values)
# True means the input logprob has been already sent to detokenizer.
@@ -918,6 +923,12 @@ class Req(ReqDllmMixin):
# 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.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
+52
View File
@@ -240,6 +240,7 @@ from sglang.srt.platforms import current_platform
from sglang.srt.plugins import load_plugins
from sglang.srt.runtime_context import get_parallel, get_server_args
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.sampling.sampling_params import TOP_K_ALL
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.session.session_controller import SessionController
from sglang.srt.speculative.dflash_utils import validate_dflash_request
@@ -2091,6 +2092,7 @@ class Scheduler(
return_logprob=recv_req.return_logprob,
top_logprobs_num=recv_req.top_logprobs_num,
token_ids_logprob=recv_req.token_ids_logprob,
return_sampling_mask=recv_req.return_sampling_mask,
stream=recv_req.stream,
lora_id=recv_req.lora_id,
session_id=recv_req.session_id,
@@ -2193,6 +2195,56 @@ class Scheduler(
self.init_req_max_new_tokens(req)
self._add_request_to_queue(req)
return
if (
req.return_sampling_mask
and self.disaggregation_mode != DisaggregationMode.NULL
and not self.disagg_metadata_buffers.enable_sampling_mask
):
error_msg = (
"return_sampling_mask with disaggregation requires "
"SGLANG_DISAGGREGATION_SAMPLING_MASK_MAX_TOKENS > 0."
)
req.set_finish_with_abort(error_msg)
self.init_req_max_new_tokens(req)
self._add_request_to_queue(req)
return
if req.return_sampling_mask and req.sampling_params.top_k == TOP_K_ALL:
error_msg = (
"return_sampling_mask requires finite top_k; top_p-only sampling "
"is valid but can return huge masks in the tail, blowing up "
"metadata, so we need a safety cap."
)
req.set_finish_with_abort(error_msg)
self.init_req_max_new_tokens(req)
self._add_request_to_queue(req)
return
if req.return_sampling_mask and not self.spec_algorithm.is_none():
# Spec workers do not emit one sampling support per accepted token, so
# the returned mask would not align 1:1 with generated tokens. Reject
# the combination instead of silently returning a misaligned mask.
error_msg = (
"return_sampling_mask is not supported with speculative decoding."
)
req.set_finish_with_abort(error_msg)
self.init_req_max_new_tokens(req)
self._add_request_to_queue(req)
return
if req.return_sampling_mask and self.server_args.sampling_backend == "ascend":
# The ascend backend samples from logits directly and never builds the
# top-k/top-p support, so it cannot produce a sampling mask.
error_msg = (
"return_sampling_mask is not supported with the ascend "
"sampling backend."
)
req.set_finish_with_abort(error_msg)
self.init_req_max_new_tokens(req)
self._add_request_to_queue(req)
return
# Handle multimodal inputs
if recv_req.mm_inputs is not None:
image_inputs = self._get_multimodal_inputs(recv_req.mm_inputs)
@@ -258,6 +258,9 @@ class SchedulerBatchResultProcessor:
logprob_pt=logprob_pt,
)
if req.return_sampling_mask:
self.add_sampling_mask_return_values(i, req, logits_output)
if (
req.return_hidden_states
and logits_output.hidden_states is not None
@@ -721,6 +724,11 @@ class SchedulerBatchResultProcessor:
logits_output=logits_output,
)
if req.return_sampling_mask:
# return_sampling_mask + speculative decoding is rejected at
# request entry, so this remains one support mask per token.
self.add_sampling_mask_return_values(i, req, logits_output)
if req.return_hidden_states and logits_output.hidden_states is not None:
# hidden_states is [bs * stride, hidden_dim], one row per emitted
# token; stride = speculative_num_draft_tokens for spec, 1 for non-spec.
@@ -831,6 +839,20 @@ class SchedulerBatchResultProcessor:
logits_output.next_token_token_ids_logprobs_idx[flat_idx]
)
def add_sampling_mask_return_values(
self,
i: int,
req: Req,
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]
)
def _handle_finish_state_updated_req(
self,
req: Req,
@@ -132,12 +132,16 @@ class SchedulerOutputStreamer:
return_indexer_topk = any(
req.return_indexer_topk for req in reqs if req is not skip_req
)
return_sampling_mask = any(
req.return_sampling_mask for req in reqs if req is not skip_req
)
acc = _GenerationStreamAccumulator(
return_logprob=return_logprob,
return_hidden_states=return_hidden_states,
return_routed_experts=return_routed_experts,
return_indexer_topk=return_indexer_topk,
return_sampling_mask=return_sampling_mask,
spec_algorithm=self.spec_algorithm,
disaggregation_mode=self.disaggregation_mode,
default_stream_interval=self.server_args.stream_interval,
@@ -250,6 +254,7 @@ class _GenerationStreamAccumulator:
return_hidden_states: bool
return_routed_experts: bool
return_indexer_topk: bool
return_sampling_mask: bool = False
spec_algorithm: Any
disaggregation_mode: DisaggregationMode
default_stream_interval: int
@@ -300,6 +305,8 @@ 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] = None
output_token_sampling_logprobs: Optional[list] = None
def __post_init__(self) -> None:
if self.return_hidden_states:
@@ -322,6 +329,9 @@ class _GenerationStreamAccumulator:
self.input_token_ids_logprobs_idx = []
self.output_token_ids_logprobs_val = []
self.output_token_ids_logprobs_idx = []
if self.return_sampling_mask:
self.output_token_sampling_mask = []
self.output_token_sampling_logprobs = []
def accept(self, *, req: Req) -> None:
if req.finished():
@@ -487,6 +497,25 @@ class _GenerationStreamAccumulator:
self.output_token_ids_logprobs_val.append([])
self.output_token_ids_logprobs_idx.append([])
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
else:
self.output_token_sampling_mask.append([])
self.output_token_sampling_logprobs.append([])
if self.return_hidden_states:
if req.return_hidden_states:
# Mirror output_ids_through_stop: spec verify steps can overshoot finished_len.
@@ -568,6 +597,8 @@ class _GenerationStreamAccumulator:
output_token_ids_logprobs_val=self.output_token_ids_logprobs_val,
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,
@@ -160,6 +160,8 @@ _INCREMENTAL_STREAMING_META_INFO_KEYS = (
"output_token_logprobs",
"output_top_logprobs",
"output_token_ids_logprobs",
"output_token_sampling_mask",
"output_token_sampling_logprobs",
)
@@ -218,6 +220,8 @@ class ReqState:
input_token_ids_logprobs_idx: List = dataclasses.field(default_factory=list)
output_token_ids_logprobs_val: List = dataclasses.field(default_factory=list)
output_token_ids_logprobs_idx: List = dataclasses.field(default_factory=list)
output_token_sampling_mask: List = dataclasses.field(default_factory=list)
output_token_sampling_logprobs: List = dataclasses.field(default_factory=list)
# For detokenized logprobs
input_token_logprobs: List[Any] = dataclasses.field(default_factory=list)
@@ -1175,6 +1179,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
logprob_start_len=obj.logprob_start_len,
top_logprobs_num=obj.top_logprobs_num,
token_ids_logprob=obj.token_ids_logprob,
return_sampling_mask=obj.return_sampling_mask,
stream=obj.stream,
rid=obj.rid,
http_worker_ipc=obj.http_worker_ipc,
@@ -1927,6 +1932,27 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
recv_obj,
i,
)
if (
isinstance(state.obj, GenerateReqInput)
and state.obj.return_sampling_mask
):
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]
)
meta_info["output_token_sampling_mask"] = (
state.output_token_sampling_mask
)
meta_info["output_token_sampling_logprobs"] = (
state.output_token_sampling_logprobs
)
meta_info["output_token_sampling_mask_length"] = len(
state.output_token_sampling_mask
)
if not isinstance(recv_obj, BatchEmbeddingOutput):
meta_info.update(
+4
View File
@@ -224,6 +224,8 @@ def get_logprob_dict_from_result(result: GenerationBatchResult) -> dict:
"next_token_top_logprobs_idx": result.logits_output.next_token_top_logprobs_idx,
"next_token_token_ids_logprobs_val": result.logits_output.next_token_token_ids_logprobs_val,
"next_token_token_ids_logprobs_idx": result.logits_output.next_token_token_ids_logprobs_idx,
"next_token_sampling_mask_idx": result.logits_output.next_token_sampling_mask_idx,
"next_token_sampling_logprobs": result.logits_output.next_token_sampling_logprobs,
"input_token_logprobs": result.logits_output.input_token_logprobs,
"input_top_logprobs_val": result.logits_output.input_top_logprobs_val,
"input_top_logprobs_idx": result.logits_output.input_top_logprobs_idx,
@@ -248,6 +250,8 @@ def get_logprob_from_pp_outputs(
next_token_token_ids_logprobs_idx=next_pp_outputs[
"next_token_token_ids_logprobs_idx"
],
next_token_sampling_mask_idx=next_pp_outputs["next_token_sampling_mask_idx"],
next_token_sampling_logprobs=next_pp_outputs["next_token_sampling_logprobs"],
input_token_logprobs=next_pp_outputs["input_token_logprobs"],
input_top_logprobs_val=next_pp_outputs["input_top_logprobs_val"],
input_top_logprobs_idx=next_pp_outputs["input_top_logprobs_idx"],
@@ -69,6 +69,10 @@ class SamplingBatchInfo:
# Used for deterministic sampling
sampling_seed: Optional[torch.Tensor] = None
# Per-request flag for returning sparse sampling support metadata.
return_sampling_masks: Optional[List[bool]] = None
sampling_mask_max_top_k: int = 0
# Device
device: str = "cuda"
@@ -137,6 +141,11 @@ class SamplingBatchInfo:
global_server_args.enable_custom_logit_processor
and any(r.custom_logit_processor for r in reqs) # check the flag first.
) # then check the requests.
return_sampling_masks = [r.return_sampling_mask for r in reqs]
sampling_mask_max_top_k = max(
(r.sampling_params.top_k for r in reqs if r.return_sampling_mask),
default=0,
)
if has_custom_logit_processor:
# Merge the same type of custom logit processors together
@@ -201,6 +210,8 @@ class SamplingBatchInfo:
custom_logit_processor=merged_custom_logit_processor,
device=device,
logit_bias=logit_bias,
return_sampling_masks=return_sampling_masks,
sampling_mask_max_top_k=sampling_mask_max_top_k,
)
ret.adjusted_from_schedule_batch(batch, vocab_size)
return ret
@@ -304,6 +315,10 @@ class SamplingBatchInfo:
if self.logit_bias is not None:
self.logit_bias = self.logit_bias[keep_indices_device]
if self.return_sampling_masks is not None:
self.return_sampling_masks = [
self.return_sampling_masks[i] for i in keep_indices
]
self.adjusted_filter_batch(keep_indices, keep_indices_device)
@@ -390,11 +405,24 @@ class SamplingBatchInfo:
# Set the flag to True if any of the two has custom logit processor
self.has_custom_logit_processor = True
self_len = len(self)
other_len = len(other)
# Merge logit bias - note this has to come before the temperatures tensor update! Otherwise will cause crashes.
# See note below on len(self) and len(other).
self.logit_bias = merge_bias_tensor(
self.logit_bias, other.logit_bias, len(self), len(other), self.device, 0.0
self.logit_bias, other.logit_bias, self_len, other_len, self.device, 0.0
)
if (
self.return_sampling_masks is not None
or other.return_sampling_masks is not None
):
self.return_sampling_masks = (
self.return_sampling_masks or [False] * self_len
) + (other.return_sampling_masks or [False] * other_len)
self.sampling_mask_max_top_k = max(
self.sampling_mask_max_top_k, other.sampling_mask_max_top_k
)
# Note: because the __len()__ operator is defined on the temperatures tensor,
# please make sure any merge operation with len(self) or len(other) is done before
@@ -302,6 +302,7 @@ class Session:
return_logprob=req.return_logprob,
top_logprobs_num=req.top_logprobs_num,
token_ids_logprob=req.token_ids_logprob,
return_sampling_mask=req.return_sampling_mask,
vocab_size=vocab_size,
eos_token_ids=eos_token_ids,
require_reasoning=req.require_reasoning,
@@ -0,0 +1,249 @@
import math
import unittest
import requests
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.test_utils import (
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
)
register_cuda_ci(est_time=240, stage="base-b", runner_config="1-gpu-small")
register_amd_ci(est_time=320, suite="stage-b-test-1-gpu-small-amd")
_MAX_NEW_TOKENS = 4
_TOP_P = 0.99
_TOP_K = 10
_SAMPLING_SEED = 1234
_SERVER_ARGS = (
"--mem-fraction-static",
"0.7",
)
_INVALID_SAMPLING_MASK_ERROR = (
"top_p-only sampling is valid but can return huge masks in the tail"
)
class SamplingMaskTestMixin:
@classmethod
def _launch_server(cls, other_args=()):
cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=(*_SERVER_ARGS, *other_args),
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
def _post_generate(
self,
sampling_params,
return_sampling_mask=True,
return_logprob=False,
top_logprobs_num=0,
):
payload = {
"text": "The capital of France is",
"sampling_params": sampling_params,
"return_sampling_mask": return_sampling_mask,
}
if return_logprob:
payload["return_logprob"] = True
payload["top_logprobs_num"] = top_logprobs_num
return requests.post(self.base_url + "/generate", json=payload, timeout=60)
def _generate_sampling_masks(self, sampling_params):
response = self._post_generate(sampling_params)
self.assertEqual(response.status_code, 200, response.text)
output = response.json()
meta_info = output["meta_info"]
output_ids = output["output_ids"]
sampling_masks = meta_info["output_token_sampling_mask"]
self.assertEqual(len(output_ids), _MAX_NEW_TOKENS)
self.assertEqual(meta_info["completion_tokens"], len(output_ids))
self.assertEqual(
meta_info["output_token_sampling_mask_length"], len(output_ids)
)
self.assertEqual(len(sampling_masks), len(output_ids))
for output_id, sampling_mask in zip(output_ids, sampling_masks):
self.assertIn(output_id, sampling_mask)
return sampling_masks
def _assert_rejects_unbounded_sampling_mask(self, sampling_params):
response = self._post_generate(sampling_params)
self.assertEqual(response.status_code, 400, response.text)
self.assertIn(_INVALID_SAMPLING_MASK_ERROR, response.text)
class TestSamplingMask(SamplingMaskTestMixin, CustomTestCase):
@classmethod
def setUpClass(cls):
cls._launch_server()
def test_generate_returns_sampling_mask(self):
top_p_sampling_masks = self._generate_sampling_masks(
{
"temperature": 1.0,
"top_k": _TOP_K,
"top_p": _TOP_P,
"max_new_tokens": _MAX_NEW_TOKENS,
"ignore_eos": True,
}
)
for sampling_mask in top_p_sampling_masks:
self.assertLessEqual(len(sampling_mask), _TOP_K)
top_k_sampling_masks = self._generate_sampling_masks(
{
"temperature": 1.0,
"top_k": _TOP_K,
"max_new_tokens": _MAX_NEW_TOKENS,
"ignore_eos": True,
}
)
for sampling_mask in top_k_sampling_masks:
self.assertEqual(len(sampling_mask), _TOP_K)
top_k_top_p_one_sampling_masks = self._generate_sampling_masks(
{
"temperature": 1.0,
"top_k": _TOP_K,
"top_p": 1.0,
"max_new_tokens": _MAX_NEW_TOKENS,
"ignore_eos": True,
}
)
for sampling_mask in top_k_top_p_one_sampling_masks:
self.assertEqual(len(sampling_mask), _TOP_K)
def test_sampling_mask_matches_topk_logprobs(self):
"""Check the returned mask and its renormalized logprobs.
We get the per-token full-vocab logprobs via ``return_logprob`` with
``top_logprobs_num == top_k``, which covers every token the mask can
contain. With ``temperature=1.0`` these are the sampler's distribution,
so ``p = exp(logprob)`` are the exact probabilities. For each token, we check:
1. the returned mask matches the nucleus reconstructed from those probs,
2. sampling_logprob == log(p[sampled] / sum(p[t] for t in mask)).
"""
top_k, top_p = _TOP_K, _TOP_P
response = self._post_generate(
{
"temperature": 1.0,
"top_k": top_k,
"top_p": top_p,
"max_new_tokens": _MAX_NEW_TOKENS,
"ignore_eos": True,
},
return_logprob=True,
top_logprobs_num=top_k,
)
self.assertEqual(response.status_code, 200, response.text)
output = response.json()
meta_info = output["meta_info"]
output_ids = output["output_ids"]
sampling_masks = meta_info["output_token_sampling_mask"]
sampling_logprobs = meta_info["output_token_sampling_logprobs"]
top_logprobs = meta_info["output_top_logprobs"] # [logprob, id, text] per token
self.assertEqual(len(sampling_masks), len(output_ids))
self.assertEqual(len(sampling_logprobs), len(output_ids))
self.assertEqual(len(top_logprobs), len(output_ids))
for output_id, mask, mask_logprob, step_top_logprobs in zip(
output_ids, sampling_masks, sampling_logprobs, top_logprobs
):
probs = {
int(tid): math.exp(logprob) for logprob, tid, _ in step_top_logprobs
}
reconstructed = []
mass_before = 0.0
for logprob, tid, _ in step_top_logprobs:
if mass_before <= top_p:
reconstructed.append(int(tid))
mass_before += math.exp(logprob)
if output_id not in reconstructed:
reconstructed.append(output_id)
# ``<= 1``: fp32 (server) and fp64 (here) cumsums may split on the
# single token straddling the top_p cut.
self.assertLessEqual(len(set(mask) ^ set(reconstructed)), 1)
support_mass = sum(probs[tid] for tid in mask)
expected_logprob = math.log(probs[output_id] / support_mass)
self.assertAlmostEqual(mask_logprob, expected_logprob, delta=1e-2)
def test_generate_rejects_unbounded_sampling_mask(self):
self._assert_rejects_unbounded_sampling_mask(
{
"temperature": 1.0,
"top_p": _TOP_P,
"max_new_tokens": _MAX_NEW_TOKENS,
"ignore_eos": True,
}
)
self._assert_rejects_unbounded_sampling_mask(
{
"temperature": 1.0,
"top_p": 1.0,
"max_new_tokens": _MAX_NEW_TOKENS,
"ignore_eos": True,
}
)
class TestSamplingMaskDeterministic(SamplingMaskTestMixin, CustomTestCase):
@classmethod
def setUpClass(cls):
# This test validates sampler/output determinism, not backend selection.
# Pin Triton so the same deterministic path runs on CUDA and ROCm CI.
cls._launch_server(
("--enable-deterministic-inference", "--attention-backend", "triton")
)
def test_return_sampling_mask_preserves_deterministic_sampling(self):
sampling_params = {
"temperature": 1.0,
"top_k": _TOP_K,
"top_p": 1.0,
"sampling_seed": _SAMPLING_SEED,
"max_new_tokens": _MAX_NEW_TOKENS,
"ignore_eos": True,
}
with_mask_response = self._post_generate(
sampling_params, return_sampling_mask=True
)
self.assertEqual(with_mask_response.status_code, 200, with_mask_response.text)
without_mask_response = self._post_generate(
sampling_params, return_sampling_mask=False
)
self.assertEqual(
without_mask_response.status_code, 200, without_mask_response.text
)
with_mask_output = with_mask_response.json()
without_mask_output = without_mask_response.json()
self.assertEqual(
with_mask_output["output_ids"], without_mask_output["output_ids"]
)
self.assertEqual(with_mask_output["text"], without_mask_output["text"])
if __name__ == "__main__":
unittest.main()
@@ -68,6 +68,8 @@ def _make_batch_str_output() -> BatchStrOutput:
output_token_ids_logprobs_val=[[], []],
output_token_ids_logprobs_idx=[[], []],
output_token_entropy_val=[0.0, 0.0],
output_token_sampling_mask=[[], []],
output_token_sampling_logprobs=[[], []],
output_hidden_states=[None, None],
routed_experts=[None, None],
indexer_topk=[None, None],
@@ -38,6 +38,7 @@ def _recv(rid, input_ids, max_new_tokens=8):
return_logprob=False,
top_logprobs_num=0,
token_ids_logprob=None,
return_sampling_mask=False,
require_reasoning=False,
return_hidden_states=False,
return_routed_experts=False,