[Sampling] Capture masks from sampler support (#36630)

Co-authored-by: ByronHsu <ByronHsu@users.noreply.github.com>
This commit is contained in:
Byron Hsu
2026-09-02 17:13:28 -07:00
committed by GitHub
co-authored by ByronHsu
parent 5ddca6819e
commit 046cdaabaa
3 changed files with 433 additions and 127 deletions
+223 -91
View File
@@ -1,6 +1,6 @@
import logging import logging
from functools import partial from functools import partial
from typing import Callable, Dict, List, Optional, Tuple from typing import Callable, Dict, List, NamedTuple, Optional, Tuple
import torch import torch
import torch.distributed as dist import torch.distributed as dist
@@ -84,6 +84,15 @@ def _trace_e2e_sampler(stage: str, **fields) -> None:
print(f"SGLANG_TRACE_SAMPLER_E2E {rank} stage={stage} {details}", flush=True) print(f"SGLANG_TRACE_SAMPLER_E2E {rank} stage={stage} {details}", flush=True)
class _SamplingMaskCapture(NamedTuple):
"""Compact post-filter weights and their original batch-row mapping."""
weights: torch.Tensor
token_ids: Optional[torch.Tensor]
selected_weight: Optional[torch.Tensor]
batch_rows: torch.Tensor
class Sampler(nn.Module): class Sampler(nn.Module):
def __init__(self): def __init__(self):
super().__init__() super().__init__()
@@ -149,6 +158,7 @@ class Sampler(nn.Module):
logits = self._preprocess_logits(logits, sampling_info) logits = self._preprocess_logits(logits, sampling_info)
_trace_e2e_sampler("preprocess_returned") _trace_e2e_sampler("preprocess_returned")
return_sampling_mask = any(sampling_info.return_sampling_masks or []) return_sampling_mask = any(sampling_info.return_sampling_masks or [])
sampling_mask_capture = None
if sampling_info.is_all_greedy: if sampling_info.is_all_greedy:
_trace_e2e_sampler("greedy_enter") _trace_e2e_sampler("greedy_enter")
@@ -162,10 +172,6 @@ class Sampler(nn.Module):
_trace_e2e_sampler( _trace_e2e_sampler(
"greedy_returned", output_shape=tuple(batch_next_token_ids.shape) "greedy_returned", output_shape=tuple(batch_next_token_ids.shape)
) )
if return_sampling_mask:
self._attach_greedy_sampling_mask_to_output(
logits_output, sampling_info, batch_next_token_ids
)
if return_logprob: if return_logprob:
original_logprobs = logprobs = torch.nn.functional.log_softmax( original_logprobs = logprobs = torch.nn.functional.log_softmax(
logits, dim=-1 logits, dim=-1
@@ -238,19 +244,13 @@ class Sampler(nn.Module):
logits[:] = torch.softmax(logits, dim=-1) logits[:] = torch.softmax(logits, dim=-1)
probs = logits probs = logits
batch_next_token_ids = self._sample_from_probs( batch_next_token_ids, sampling_mask_capture = self._sample_from_probs(
probs, sampling_info, positions, simple_sampling_case probs,
sampling_info,
positions,
simple_sampling_case,
return_sampling_mask=return_sampling_mask,
) )
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: if return_logprob and not SGLANG_RETURN_ORIGINAL_LOGPROB:
logprobs = ( logprobs = (
logprobs_via_logsoftmax_kernel logprobs_via_logsoftmax_kernel
@@ -274,6 +274,25 @@ class Sampler(nn.Module):
self._sync_token_ids_across_tp(batch_next_token_ids, sampling_info) self._sync_token_ids_across_tp(batch_next_token_ids, sampling_info)
_trace_e2e_sampler("token_sync_returned") _trace_e2e_sampler("token_sync_returned")
if return_sampling_mask:
if sampling_info.is_all_greedy:
self._attach_greedy_sampling_mask_to_output(
logits_output, sampling_info, batch_next_token_ids
)
else:
assert sampling_mask_capture is not None
if SYNC_TOKEN_IDS_ACROSS_TP or sampling_info.grammars:
# Token synchronization can replace the producer-selected token.
sampling_mask_capture = sampling_mask_capture._replace(
selected_weight=None
)
self._attach_sampling_mask_to_output(
logits_output,
sampling_info,
batch_next_token_ids,
sampling_mask_capture,
)
_trace_e2e_sampler("forward_returned") _trace_e2e_sampler("forward_returned")
return batch_next_token_ids return batch_next_token_ids
@@ -283,18 +302,57 @@ class Sampler(nn.Module):
sampling_info: SamplingBatchInfo, sampling_info: SamplingBatchInfo,
positions: torch.Tensor, positions: torch.Tensor,
simple_sampling_case: bool, simple_sampling_case: bool,
) -> torch.Tensor: *,
return_sampling_mask: bool = False,
) -> Tuple[torch.Tensor, Optional[_SamplingMaskCapture]]:
"""Sample from probability distribution (after softmax). """Sample from probability distribution (after softmax).
Used for standard sampling with flashinfer/pytorch backends. Used for standard sampling with flashinfer/pytorch backends.
Handles both simple (direct multinomial) and complex (top-k/top-p/min-p) cases. Handles both simple (direct multinomial) and complex (top-k/top-p/min-p) cases.
Capture work is performed only when return_sampling_mask is enabled.
""" """
sampling_mask_capture = None
capture_rows = None
capture_all_rows = False
if return_sampling_mask:
capture_rows_list = [
i
for i, should_return in enumerate(
sampling_info.return_sampling_masks or []
)
if should_return
]
if not capture_rows_list:
raise RuntimeError(
"Sampling-mask capture requested without any opted-in batch rows."
)
capture_rows = torch.tensor(
capture_rows_list, device=probs.device, dtype=torch.long
)
capture_all_rows = capture_rows_list == list(range(probs.shape[0]))
def select_capture_rows(tensor: torch.Tensor) -> torch.Tensor:
assert capture_rows is not None
return tensor if capture_all_rows else tensor.index_select(0, capture_rows)
if simple_sampling_case: if simple_sampling_case:
batch_next_token_ids = sampling_from_probs_torch( batch_next_token_ids = sampling_from_probs_torch(
probs, probs,
sampling_seed=sampling_info.sampling_seed, sampling_seed=sampling_info.sampling_seed,
positions=positions, positions=positions,
) )
if return_sampling_mask:
capture_probs = select_capture_rows(probs)
capture_tokens = select_capture_rows(batch_next_token_ids)
selected_weight = torch.gather(
capture_probs, 1, capture_tokens.long().view(-1, 1)
).squeeze(1)
sampling_mask_capture = _SamplingMaskCapture(
weights=capture_probs,
token_ids=None,
selected_weight=selected_weight,
batch_rows=capture_rows,
)
else: else:
backend = get_exec().kernel.sampling_backend backend = get_exec().kernel.sampling_backend
if backend == "flashinfer": if backend == "flashinfer":
@@ -307,6 +365,27 @@ class Sampler(nn.Module):
batch_next_token_ids = min_p_sampling_from_probs( batch_next_token_ids = min_p_sampling_from_probs(
probs, sampling_info.min_ps probs, sampling_info.min_ps
) )
if return_sampling_mask:
capture_probs = select_capture_rows(probs)
capture_min_ps = select_capture_rows(sampling_info.min_ps)
capture_tokens = select_capture_rows(batch_next_token_ids)
min_p_thresholds = (
capture_probs.max(dim=-1).values * capture_min_ps
)
filtered_probs = capture_probs.masked_fill(
capture_probs < min_p_thresholds.view(-1, 1), 0
)
selected_weight = torch.gather(
filtered_probs,
1,
capture_tokens.long().view(-1, 1),
).squeeze(1)
sampling_mask_capture = _SamplingMaskCapture(
weights=filtered_probs,
token_ids=None,
selected_weight=selected_weight,
batch_rows=capture_rows,
)
else: else:
batch_next_token_ids = top_k_top_p_sampling_from_probs( batch_next_token_ids = top_k_top_p_sampling_from_probs(
probs.contiguous(), probs.contiguous(),
@@ -314,9 +393,42 @@ class Sampler(nn.Module):
sampling_info.top_ps, sampling_info.top_ps,
filter_apply_order="joint", filter_apply_order="joint",
) )
if return_sampling_mask:
# Correctness invariant: the fused joint sampler and these
# separate renormalization primitives must share cutoff,
# tie, and joint-support semantics so captured positive
# support exactly describes the sampler's action space.
capture_probs = select_capture_rows(probs)
capture_top_ks = select_capture_rows(sampling_info.top_ks)
capture_top_ps = select_capture_rows(sampling_info.top_ps)
capture_tokens = select_capture_rows(batch_next_token_ids)
filtered_probs = capture_probs
if sampling_info.need_top_k_sampling:
filtered_probs = top_k_renorm_prob(
capture_probs, capture_top_ks
)
if sampling_info.need_top_p_sampling:
top_p_probs = top_p_renorm_prob(
capture_probs, capture_top_ps
)
if filtered_probs is capture_probs:
filtered_probs = top_p_probs
else:
filtered_probs.masked_fill_(top_p_probs <= 0, 0)
selected_weight = torch.gather(
filtered_probs,
1,
capture_tokens.long().view(-1, 1),
).squeeze(1)
sampling_mask_capture = _SamplingMaskCapture(
weights=filtered_probs,
token_ids=None,
selected_weight=selected_weight,
batch_rows=capture_rows,
)
elif backend == "pytorch": elif backend == "pytorch":
# A slower fallback implementation with torch native operations. # A slower fallback implementation with torch native operations.
batch_next_token_ids = top_k_top_p_min_p_sampling_from_probs_torch( sample_result = top_k_top_p_min_p_sampling_from_probs_torch(
probs, probs,
sampling_info.top_ks, sampling_info.top_ks,
sampling_info.top_ps, sampling_info.top_ps,
@@ -324,39 +436,26 @@ class Sampler(nn.Module):
sampling_info.need_min_p_sampling, sampling_info.need_min_p_sampling,
sampling_info.sampling_seed, sampling_info.sampling_seed,
positions, positions,
return_filtered_probs=return_sampling_mask,
) )
if return_sampling_mask:
(
batch_next_token_ids,
filtered_probs,
token_ids,
selected_weight,
) = sample_result
sampling_mask_capture = _SamplingMaskCapture(
weights=select_capture_rows(filtered_probs),
token_ids=select_capture_rows(token_ids),
selected_weight=select_capture_rows(selected_weight),
batch_rows=capture_rows,
)
else:
batch_next_token_ids = sample_result
else: else:
raise ValueError(f"Invalid sampling backend: {backend}") raise ValueError(f"Invalid sampling backend: {backend}")
return batch_next_token_ids return batch_next_token_ids, sampling_mask_capture
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( def _attach_greedy_sampling_mask_to_output(
self, self,
@@ -382,63 +481,84 @@ class Sampler(nn.Module):
logits_output: LogitsProcessorOutput, logits_output: LogitsProcessorOutput,
sampling_info: SamplingBatchInfo, sampling_info: SamplingBatchInfo,
batch_next_token_ids: torch.Tensor, batch_next_token_ids: torch.Tensor,
sampling_mask_data: Tuple[ sampling_mask_capture: _SamplingMaskCapture,
torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor
],
) -> None: ) -> None:
probs_idx, probs_sort, keep_mask, probs = sampling_mask_data
return_sampling_masks = sampling_info.return_sampling_masks or [] return_sampling_masks = sampling_info.return_sampling_masks or []
if not return_sampling_masks: if not return_sampling_masks:
logits_output.next_token_sampling_mask_idx = [] logits_output.next_token_sampling_mask_idx = []
logits_output.next_token_sampling_logprobs = [] logits_output.next_token_sampling_logprobs = []
return return
sampled_tokens = batch_next_token_ids.view(-1, 1) requested_rows_list = [
sampled_matches_all = probs_idx == sampled_tokens i for i, should_return in enumerate(return_sampling_masks) if should_return
sampled_in_idx = sampled_matches_all.any(dim=-1) ]
requested_rows = sampling_mask_capture.batch_rows
# The sampler is the source of truth for the rollout action space. If a weights = sampling_mask_capture.weights
# backend/numeric edge chooses a token just outside the reconstructed token_ids = sampling_mask_capture.token_ids
# prefix, include that sampled token so training can replay a support selected_weight = sampling_mask_capture.selected_weight
# that contained the rollout action. sampled_tokens = batch_next_token_ids.index_select(0, requested_rows).view(
effective_keep_mask = keep_mask | sampled_matches_all -1, 1
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)
) )
if token_ids is None:
support_token_ids = (
torch.arange(
weights.shape[-1], device=weights.device, dtype=torch.int32
)
.view(1, -1)
.expand_as(weights)
)
selected_from_weights = torch.gather(
weights, 1, sampled_tokens.long()
).squeeze(1)
sampled_in_capture = selected_from_weights > 0
else:
sampled_matches = token_ids == sampled_tokens.to(token_ids.dtype)
sampled_in_capture = sampled_matches.any(dim=-1)
selected_positions = sampled_matches.to(torch.int32).argmax(
dim=-1, keepdim=True
)
selected_from_weights = torch.gather(
weights, 1, selected_positions
).squeeze(1)
support_token_ids = token_ids
flat_rows, flat_cols = effective_keep_mask.nonzero(as_tuple=True) if selected_weight is None:
flat_ids = probs_idx[flat_rows, flat_cols].to(torch.int32) selected_weight = selected_from_weights
mask_lengths = effective_keep_mask.sum(dim=-1, dtype=torch.int32)
support = weights > 0
support_mass = weights.sum(dim=-1, dtype=torch.float32)
selected_weight = selected_weight.float()
selected_logprobs = torch.log(selected_weight / support_mass)
valid = (
sampled_in_capture
& (selected_weight > 0)
& (support_mass > 0)
& torch.isfinite(selected_logprobs)
)
if not bool(torch.all(valid).item()):
invalid_rows = (~valid).nonzero(as_tuple=True)[0].cpu().tolist()
raise RuntimeError(
"Sampled token is outside captured positive sampling support "
f"for batch rows {invalid_rows}."
)
flat_rows, flat_cols = support.nonzero(as_tuple=True)
flat_ids = support_token_ids[flat_rows, flat_cols].to(torch.int32)
mask_lengths = support.sum(dim=-1, dtype=torch.int32)
flat_ids_cpu = flat_ids.cpu().tolist() flat_ids_cpu = flat_ids.cpu().tolist()
mask_lengths_cpu = mask_lengths.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() selected_logprobs_cpu = selected_logprobs.cpu().tolist()
masks = [] masks = [None] * len(return_sampling_masks)
logprobs = [] logprobs = [None] * len(return_sampling_masks)
cursor = 0 cursor = 0
for i, should_return in enumerate(return_sampling_masks): for capture_row, batch_row in enumerate(requested_rows_list):
mask_len = int(mask_lengths_cpu[i]) mask_len = int(mask_lengths_cpu[capture_row])
row_ids = flat_ids_cpu[cursor : cursor + mask_len] row_ids = flat_ids_cpu[cursor : cursor + mask_len]
cursor += mask_len cursor += mask_len
if not sampled_in_idx_cpu[i]: masks[batch_row] = row_ids
row_ids.append(int(sampled_tokens_cpu[i])) logprobs[batch_row] = float(selected_logprobs_cpu[capture_row])
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_mask_idx = masks
logits_output.next_token_sampling_logprobs = logprobs logits_output.next_token_sampling_logprobs = logprobs
@@ -602,11 +722,17 @@ def top_k_top_p_min_p_sampling_from_probs_torch(
need_min_p_sampling: bool, need_min_p_sampling: bool,
sampling_seed: Optional[torch.Tensor], sampling_seed: Optional[torch.Tensor],
positions: torch.Tensor, positions: torch.Tensor,
*,
return_filtered_probs: bool = False,
): ):
""" """
A top-k, top-p and min-p sampling implementation with native pytorch operations. A top-k, top-p and min-p sampling implementation with native pytorch operations.
When sampling_seed is not None, deterministic inference will be enabled, it will sample When sampling_seed is not None, deterministic inference will be enabled, it will sample
with the sampling_seed of each request. with the sampling_seed of each request.
By default, returns only sampled token IDs. With return_filtered_probs=True,
also returns the actual filtered weights, their token-ID permutation, and
the selected weights.
""" """
probs_sort, probs_idx = probs.sort(dim=-1, descending=True) probs_sort, probs_idx = probs.sort(dim=-1, descending=True)
probs_sum = torch.cumsum(probs_sort, dim=-1) probs_sum = torch.cumsum(probs_sort, dim=-1)
@@ -631,14 +757,20 @@ def top_k_top_p_min_p_sampling_from_probs_torch(
# apply log to get logprobs. Therefore, we cannot use log_softmax directly. # apply log to get logprobs. Therefore, we cannot use log_softmax directly.
# For now, we use log to the modified probs to get logprobs, but for numerical # For now, we use log to the modified probs to get logprobs, but for numerical
# stability, we'd better come up with a solution to use log_softmax. # stability, we'd better come up with a solution to use log_softmax.
logprobs = probs_sort.to(torch.float64) # Using float64 for numerical stability logprobs = probs_sort.to(torch.float64, copy=return_filtered_probs)
del probs_sort if not return_filtered_probs:
del probs_sort
logprobs.log_() logprobs.log_()
sampled_index = multinomial_with_seed(logprobs, sampling_seed, positions) sampled_index = multinomial_with_seed(logprobs, sampling_seed, positions)
if return_filtered_probs:
selected_weight = torch.gather(probs_sort, 1, sampled_index).view(-1)
# int32 range is enough to represent the token ids # int32 range is enough to represent the token ids
probs_idx = probs_idx.to(torch.int32) probs_idx = probs_idx.to(torch.int32)
batch_next_token_ids = torch.gather(probs_idx, dim=1, index=sampled_index).view(-1) batch_next_token_ids = torch.gather(probs_idx, dim=1, index=sampled_index).view(-1)
if return_filtered_probs:
return batch_next_token_ids, probs_sort, probs_idx, selected_weight
return batch_next_token_ids return batch_next_token_ids
@@ -76,7 +76,6 @@ class SamplingBatchInfo:
# Per-request flag for returning sparse sampling support metadata. # Per-request flag for returning sparse sampling support metadata.
return_sampling_masks: Optional[List[bool]] = None return_sampling_masks: Optional[List[bool]] = None
sampling_mask_max_top_k: int = 0
# Device # Device
device: str = "cuda" device: str = "cuda"
@@ -146,10 +145,6 @@ class SamplingBatchInfo:
and any(r.custom_logit_processor for r in reqs) # check the flag first. and any(r.custom_logit_processor for r in reqs) # check the flag first.
) # then check the requests. ) # then check the requests.
return_sampling_masks = [r.return_sampling_mask for r in reqs] 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: if has_custom_logit_processor:
# Merge the same type of custom logit processors together # Merge the same type of custom logit processors together
@@ -215,7 +210,6 @@ class SamplingBatchInfo:
device=device, device=device,
logit_bias=logit_bias, logit_bias=logit_bias,
return_sampling_masks=return_sampling_masks, return_sampling_masks=return_sampling_masks,
sampling_mask_max_top_k=sampling_mask_max_top_k,
) )
ret.adjusted_from_schedule_batch(batch, vocab_size) ret.adjusted_from_schedule_batch(batch, vocab_size)
return ret return ret
@@ -445,9 +439,6 @@ class SamplingBatchInfo:
self.return_sampling_masks = ( self.return_sampling_masks = (
self.return_sampling_masks or [False] * self_len self.return_sampling_masks or [False] * self_len
) + (other.return_sampling_masks or [False] * other_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, # 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 # please make sure any merge operation with len(self) or len(other) is done before
+210 -27
View File
@@ -1,9 +1,15 @@
import math import math
import unittest import unittest
from types import SimpleNamespace
from unittest.mock import patch
import requests import requests
import torch
from sglang.srt.utils import kill_process_tree from sglang.srt.layers import sampler as sampler_module
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.sampler import Sampler
from sglang.srt.utils import is_hip, kill_process_tree
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.test_utils import ( from sglang.test.test_utils import (
DEFAULT_SMALL_MODEL_NAME_FOR_TEST, DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
@@ -19,6 +25,7 @@ register_amd_ci(est_time=320, suite="stage-b-test-1-gpu-small-amd")
_MAX_NEW_TOKENS = 4 _MAX_NEW_TOKENS = 4
_TOP_P = 0.99 _TOP_P = 0.99
_TOP_K = 10 _TOP_K = 10
_TOP_LOGPROBS_NUM = 128
_SAMPLING_SEED = 1234 _SAMPLING_SEED = 1234
_SERVER_ARGS = ( _SERVER_ARGS = (
"--mem-fraction-static", "--mem-fraction-static",
@@ -29,6 +36,180 @@ _INVALID_SAMPLING_MASK_ERROR = (
) )
class TestSamplingMaskCapture(CustomTestCase):
def setUp(self):
self.sampler = Sampler.__new__(Sampler)
torch.nn.Module.__init__(self.sampler)
@unittest.skipIf(is_hip(), "FlashInfer is not available on ROCm")
def test_flashinfer_joint_cutoff_ties_match_capture(self):
batch_size = 256
top_k = 2
top_p = 0.45
base_probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device="cuda")
probs = base_probs.repeat(batch_size, 1)
# Derive the threshold-based joint support independently. Both filters
# cut at 0.2, so the tied entries must survive even though this yields
# more support entries than top_k.
sorted_probs = base_probs[0].sort(descending=True).values
top_k_cutoff = sorted_probs[top_k - 1]
mass_before = sorted_probs.cumsum(dim=-1) - sorted_probs
top_p_cutoff = sorted_probs[mass_before <= top_p][-1]
expected_support = (base_probs[0] >= top_k_cutoff) & (
base_probs[0] >= top_p_cutoff
)
expected_ids = expected_support.nonzero(as_tuple=True)[0].tolist()
self.assertEqual(expected_ids, [0, 1, 2])
sampling_info = SimpleNamespace(
sampling_seed=None,
need_top_k_sampling=True,
need_top_p_sampling=True,
need_min_p_sampling=False,
top_ks=torch.full((batch_size,), top_k, dtype=torch.int32, device="cuda"),
top_ps=torch.full((batch_size,), top_p, device="cuda"),
min_ps=torch.zeros(batch_size, device="cuda"),
return_sampling_masks=[True] * batch_size,
)
with patch(
"sglang.srt.layers.sampler.get_exec",
return_value=SimpleNamespace(
kernel=SimpleNamespace(sampling_backend="flashinfer")
),
):
sampled, capture = self.sampler._sample_from_probs(
probs,
sampling_info,
positions=torch.zeros(batch_size, dtype=torch.int64, device="cuda"),
simple_sampling_case=False,
return_sampling_mask=True,
)
self.assertIsNotNone(capture)
self.assertEqual(capture.batch_rows.cpu().tolist(), list(range(batch_size)))
actual_support = capture.weights > 0
self.assertTrue(
torch.equal(actual_support, expected_support.expand_as(actual_support))
)
self.assertGreater(int(actual_support[0].sum().item()), top_k)
self.assertTrue(
bool(actual_support.gather(1, sampled.view(-1, 1)).all().item())
)
@unittest.skipIf(is_hip(), "FlashInfer is not available on ROCm")
def test_flashinfer_capture_only_materializes_requested_rows(self):
batch_size = 4
top_k = 2
top_p = 0.45
requested_rows = [1, 3]
probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device="cuda").repeat(
batch_size, 1
)
sampling_info = SimpleNamespace(
sampling_seed=None,
need_top_k_sampling=True,
need_top_p_sampling=True,
need_min_p_sampling=False,
top_ks=torch.full((batch_size,), top_k, dtype=torch.int32, device="cuda"),
top_ps=torch.full((batch_size,), top_p, device="cuda"),
min_ps=torch.zeros(batch_size, device="cuda"),
return_sampling_masks=[False, True, False, True],
)
top_k_renorm = sampler_module.top_k_renorm_prob
top_p_renorm = sampler_module.top_p_renorm_prob
with (
patch(
"sglang.srt.layers.sampler.get_exec",
return_value=SimpleNamespace(
kernel=SimpleNamespace(sampling_backend="flashinfer")
),
),
patch(
"sglang.srt.layers.sampler.top_k_renorm_prob",
wraps=top_k_renorm,
) as top_k_mock,
patch(
"sglang.srt.layers.sampler.top_p_renorm_prob",
wraps=top_p_renorm,
) as top_p_mock,
):
sampled, capture = self.sampler._sample_from_probs(
probs,
sampling_info,
positions=torch.zeros(batch_size, dtype=torch.int64, device="cuda"),
simple_sampling_case=False,
return_sampling_mask=True,
)
self.assertIsNotNone(capture)
self.assertEqual(capture.batch_rows.cpu().tolist(), requested_rows)
self.assertEqual(tuple(capture.weights.shape), (len(requested_rows), 5))
self.assertEqual(tuple(top_k_mock.call_args.args[0].shape), (2, 5))
self.assertEqual(tuple(top_p_mock.call_args.args[0].shape), (2, 5))
output = LogitsProcessorOutput(next_token_logits=None)
self.sampler._attach_sampling_mask_to_output(
output, sampling_info, sampled, capture
)
self.assertIsNone(output.next_token_sampling_mask_idx[0])
self.assertEqual(set(output.next_token_sampling_mask_idx[1]), {0, 1, 2})
self.assertIsNone(output.next_token_sampling_mask_idx[2])
self.assertEqual(set(output.next_token_sampling_mask_idx[3]), {0, 1, 2})
self.assertIsNone(output.next_token_sampling_logprobs[0])
self.assertIsNotNone(output.next_token_sampling_logprobs[1])
self.assertIsNone(output.next_token_sampling_logprobs[2])
self.assertIsNotNone(output.next_token_sampling_logprobs[3])
def test_pytorch_capture_compacts_requested_rows(self):
batch_size = 4
requested_rows = [1, 3]
probs = torch.tensor([[0.4, 0.2, 0.2, 0.1, 0.1]], device="cuda").repeat(
batch_size, 1
)
sampling_info = SimpleNamespace(
sampling_seed=None,
need_top_k_sampling=True,
need_top_p_sampling=True,
need_min_p_sampling=False,
top_ks=torch.full((batch_size,), 2, dtype=torch.int32, device="cuda"),
top_ps=torch.full((batch_size,), 0.45, device="cuda"),
min_ps=torch.zeros(batch_size, device="cuda"),
return_sampling_masks=[False, True, False, True],
)
with patch(
"sglang.srt.layers.sampler.get_exec",
return_value=SimpleNamespace(
kernel=SimpleNamespace(sampling_backend="pytorch")
),
):
sampled, capture = self.sampler._sample_from_probs(
probs,
sampling_info,
positions=torch.zeros(batch_size, dtype=torch.int64, device="cuda"),
simple_sampling_case=False,
return_sampling_mask=True,
)
self.assertIsNotNone(capture)
self.assertEqual(capture.batch_rows.cpu().tolist(), requested_rows)
self.assertEqual(tuple(capture.weights.shape), (len(requested_rows), 5))
self.assertEqual(tuple(capture.token_ids.shape), (len(requested_rows), 5))
output = LogitsProcessorOutput(next_token_logits=None)
self.sampler._attach_sampling_mask_to_output(
output, sampling_info, sampled, capture
)
for batch_row in requested_rows:
self.assertIn(
int(sampled[batch_row]),
output.next_token_sampling_mask_idx[batch_row],
)
self.assertIsNotNone(output.next_token_sampling_logprobs[batch_row])
self.assertIsNone(output.next_token_sampling_mask_idx[0])
self.assertIsNone(output.next_token_sampling_mask_idx[2])
class SamplingMaskTestMixin: class SamplingMaskTestMixin:
@classmethod @classmethod
def _launch_server(cls, other_args=()): def _launch_server(cls, other_args=()):
@@ -79,6 +260,7 @@ class SamplingMaskTestMixin:
self.assertEqual(len(sampling_masks), len(output_ids)) self.assertEqual(len(sampling_masks), len(output_ids))
for output_id, sampling_mask in zip(output_ids, sampling_masks): for output_id, sampling_mask in zip(output_ids, sampling_masks):
self.assertIn(output_id, sampling_mask) self.assertIn(output_id, sampling_mask)
self.assertEqual(len(sampling_mask), len(set(sampling_mask)))
return sampling_masks return sampling_masks
def _assert_rejects_unbounded_sampling_mask(self, sampling_params): def _assert_rejects_unbounded_sampling_mask(self, sampling_params):
@@ -88,6 +270,8 @@ class SamplingMaskTestMixin:
class TestSamplingMask(SamplingMaskTestMixin, CustomTestCase): class TestSamplingMask(SamplingMaskTestMixin, CustomTestCase):
_sampling_backend = "flashinfer"
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
cls._launch_server() cls._launch_server()
@@ -102,12 +286,8 @@ class TestSamplingMask(SamplingMaskTestMixin, CustomTestCase):
"ignore_eos": True, "ignore_eos": True,
} }
) )
# The mask keeps at most top_k tokens, plus possibly the actually
# sampled token when the sampling kernel picks one just outside the
# mask's topk reconstruction (fp cumsum divergence); see
# Sampler._attach_sampling_mask_to_output.
for sampling_mask in top_p_sampling_masks: for sampling_mask in top_p_sampling_masks:
self.assertLessEqual(len(sampling_mask), _TOP_K + 1) self.assertGreater(len(sampling_mask), 0)
top_k_sampling_masks = self._generate_sampling_masks( top_k_sampling_masks = self._generate_sampling_masks(
{ {
@@ -118,7 +298,7 @@ class TestSamplingMask(SamplingMaskTestMixin, CustomTestCase):
} }
) )
for sampling_mask in top_k_sampling_masks: for sampling_mask in top_k_sampling_masks:
self.assertIn(len(sampling_mask), (_TOP_K, _TOP_K + 1)) self.assertGreaterEqual(len(sampling_mask), _TOP_K)
top_k_top_p_one_sampling_masks = self._generate_sampling_masks( top_k_top_p_one_sampling_masks = self._generate_sampling_masks(
{ {
@@ -130,18 +310,19 @@ class TestSamplingMask(SamplingMaskTestMixin, CustomTestCase):
} }
) )
for sampling_mask in top_k_top_p_one_sampling_masks: for sampling_mask in top_k_top_p_one_sampling_masks:
self.assertIn(len(sampling_mask), (_TOP_K, _TOP_K + 1)) self.assertGreaterEqual(len(sampling_mask), _TOP_K)
def test_sampling_mask_matches_topk_logprobs(self): def test_sampling_mask_matches_topk_logprobs(self):
"""Check the returned mask and its renormalized logprobs. """Check the returned mask and its renormalized logprobs.
We get the per-token full-vocab logprobs via ``return_logprob`` with We get a wide prefix of full-vocab logprobs via ``return_logprob`` so
``top_logprobs_num == top_k``, which covers every token the mask can cutoff ties that extend beyond ``top_k`` are visible. With
contain. With ``temperature=1.0`` these are the sampler's distribution, ``temperature=1.0`` these are the sampler's distribution, so
so ``p = exp(logprob)`` are the exact probabilities. For each token, we check: ``p = exp(logprob)`` are the exact probabilities. For each token, we check:
1. the returned mask matches the nucleus reconstructed from those probs, 1. the sampled token is in the returned top-k-bounded mask,
2. sampling_logprob == log(p[sampled] / sum(p[t] for t in mask)). 2. every mask token is present in the returned top logprobs,
3. sampling_logprob == log(p[sampled] / sum(p[t] for t in mask)).
""" """
top_k, top_p = _TOP_K, _TOP_P top_k, top_p = _TOP_K, _TOP_P
response = self._post_generate( response = self._post_generate(
@@ -153,7 +334,7 @@ class TestSamplingMask(SamplingMaskTestMixin, CustomTestCase):
"ignore_eos": True, "ignore_eos": True,
}, },
return_logprob=True, return_logprob=True,
top_logprobs_num=top_k, top_logprobs_num=_TOP_LOGPROBS_NUM,
) )
self.assertEqual(response.status_code, 200, response.text) self.assertEqual(response.status_code, 200, response.text)
@@ -175,19 +356,13 @@ class TestSamplingMask(SamplingMaskTestMixin, CustomTestCase):
int(tid): math.exp(logprob) for logprob, tid, _ in step_top_logprobs int(tid): math.exp(logprob) for logprob, tid, _ in step_top_logprobs
} }
reconstructed = [] mask_set = set(mask)
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) self.assertIn(output_id, mask_set)
self.assertLessEqual(len(mask_set), top_k)
self.assertTrue(mask_set.issubset(probs))
support_mass = sum(probs[token_id] for token_id in mask_set)
expected_logprob = math.log(probs[output_id] / support_mass) expected_logprob = math.log(probs[output_id] / support_mass)
self.assertAlmostEqual(mask_logprob, expected_logprob, delta=1e-2) self.assertAlmostEqual(mask_logprob, expected_logprob, delta=1e-2)
@@ -280,5 +455,13 @@ class TestSamplingMaskDeterministic(SamplingMaskTestMixin, CustomTestCase):
self.assertEqual(with_mask_output["text"], without_mask_output["text"]) self.assertEqual(with_mask_output["text"], without_mask_output["text"])
class TestSamplingMaskPytorch(TestSamplingMask):
_sampling_backend = "pytorch"
@classmethod
def setUpClass(cls):
cls._launch_server(("--sampling-backend", "pytorch"))
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()