Make the functions in logits_processor.py and sampler.py more modular (#17885)

This commit is contained in:
Lianmin Zheng
2026-01-28 16:24:23 -08:00
committed by GitHub
parent 0998de088b
commit d4180815a4
2 changed files with 433 additions and 296 deletions

View File

@@ -268,114 +268,15 @@ class LogitsProcessor(nn.Module):
self.final_logit_softcapping = None
self.return_full_logits = return_full_logits
self.multi_item_delimiter = (
get_global_server_args().multi_item_scoring_delimiter
)
# enable chunked logprobs processing
self.enable_logprobs_chunk = envs.SGLANG_ENABLE_LOGITS_PROCESSER_CHUNK.get()
# chunk size for logprobs processing
self.logprobs_chunk_size = envs.SGLANG_LOGITS_PROCESSER_CHUNK_SIZE.get()
def compute_logprobs_for_multi_item_scoring(
self,
input_ids,
hidden_states,
lm_head: VocabParallelEmbedding,
logits_metadata: Union[LogitsMetadata, ForwardBatch],
delimiter_token: int,
):
"""
Compute logprobs for multi-item scoring using delimiter-based token extraction.
This method is designed for scenarios where you want to score multiple items/candidates
against a single query by combining them into one sequence separated by delimiters.
Sequence format: Query<delimiter>Item1<delimiter>Item2<delimiter>...
Scoring positions: Extracts logprobs at positions before each <delimiter>
Args:
input_ids (torch.Tensor): Input token IDs containing query and items separated by delimiters.
Shape: [total_sequence_length] for single request or [batch_total_length] for batch.
hidden_states (torch.Tensor): Hidden states from the model.
Shape: [sequence_length, hidden_dim].
lm_head (VocabParallelEmbedding): Language model head for computing logits.
logits_metadata (Union[LogitsMetadata, ForwardBatch]): Metadata containing batch info
and token ID specifications for logprob extraction.
delimiter_token (int): Token ID used as delimiter between query and items.
Returns:
LogitsProcessorOutput: Contains:
- next_token_logits: None (not needed for scoring-only requests)
- input_token_logprobs: Logprobs of delimiter tokens at scoring positions
- input_top_logprobs_val: Top-k logprobs at delimiter positions (if requested)
- input_top_logprobs_idx: Top-k token indices at delimiter positions (if requested)
- input_token_ids_logprobs_val: Logprobs for user-requested token IDs (if any)
- input_token_ids_logprobs_idx: Indices for user-requested token IDs (if any)
"""
multi_item_indices = (input_ids == delimiter_token).nonzero(as_tuple=True)[
0
] - 1
# Extract hidden states at delimiter positions for multi-item scoring
sliced_hidden = hidden_states[multi_item_indices]
sliced_logits = self._get_logits(sliced_hidden, lm_head, logits_metadata)
sliced_logprobs = torch.nn.functional.log_softmax(sliced_logits, dim=-1)
# Initialize return values
input_token_ids_logprobs_val = []
input_token_ids_logprobs_idx = []
input_top_logprobs_val = None
input_top_logprobs_idx = None
# Recalculate extend_logprob_pruned_lens_cpu to match delimiter counts per request
# Original contains sequence lengths, but we need delimiter counts for sliced_logprobs
if (
logits_metadata.token_ids_logprobs
or logits_metadata.extend_return_top_logprob
):
logits_metadata.extend_logprob_pruned_lens_cpu = []
if logits_metadata.extend_seq_lens_cpu is not None:
# Multi-request batch: count delimiters per request
input_pt = 0
for req_seq_len in logits_metadata.extend_seq_lens_cpu:
req_input_ids = input_ids[input_pt : input_pt + req_seq_len]
delimiter_count = (req_input_ids == delimiter_token).sum().item()
logits_metadata.extend_logprob_pruned_lens_cpu.append(
delimiter_count
)
input_pt += req_seq_len
else:
# Single request case: one request gets all delimiters
total_delimiters = (input_ids == delimiter_token).sum().item()
logits_metadata.extend_logprob_pruned_lens_cpu = [total_delimiters]
# Get the logprobs of specified token ids
if logits_metadata.extend_token_ids_logprob:
(
input_token_ids_logprobs_val,
input_token_ids_logprobs_idx,
) = get_token_ids_logprobs_prefill(
sliced_logprobs, logits_metadata, delay_cpu_copy=True
)
# Get the logprob of top-k tokens
if logits_metadata.extend_return_top_logprob:
(
input_top_logprobs_val,
input_top_logprobs_idx,
) = get_top_logprobs_prefill(sliced_logprobs, logits_metadata)
# For input_token_logprobs, use delimiter token logprobs
input_token_logprobs = sliced_logprobs[:, delimiter_token]
return LogitsProcessorOutput(
next_token_logits=None, # Multi-item scoring doesn't need next token logits
input_token_logprobs=input_token_logprobs,
input_top_logprobs_val=input_top_logprobs_val,
input_top_logprobs_idx=input_top_logprobs_idx,
input_token_ids_logprobs_val=input_token_ids_logprobs_val,
input_token_ids_logprobs_idx=input_token_ids_logprobs_idx,
)
def forward(
self,
input_ids,
@@ -388,23 +289,116 @@ class LogitsProcessor(nn.Module):
if isinstance(logits_metadata, ForwardBatch):
logits_metadata = LogitsMetadata.from_forward_batch(logits_metadata)
# Check if multi-item scoring is enabled via server args (only for prefill-only requests)
multi_item_delimiter = get_global_server_args().multi_item_scoring_delimiter
if multi_item_delimiter is not None and logits_metadata.is_prefill_only:
# Multi-item scoring only for prefill-only requests.
if self.multi_item_delimiter is not None and logits_metadata.is_prefill_only:
return self.compute_logprobs_for_multi_item_scoring(
input_ids, hidden_states, lm_head, logits_metadata, multi_item_delimiter
input_ids,
hidden_states,
lm_head,
logits_metadata,
self.multi_item_delimiter,
)
# Diffusion LLM only.
if logits_metadata.forward_mode.is_dllm_extend():
assert self.return_full_logits
full_logits = self._get_logits(hidden_states, lm_head, logits_metadata)
return LogitsProcessorOutput(
full_logits=full_logits,
next_token_logits=None,
)
return self._get_dllm_logits(hidden_states, lm_head, logits_metadata)
# Get the last hidden states and last logits for the next token prediction
(
pruned_states,
pruned_states_before_norm,
aux_pruned_states,
sample_indices,
input_logprob_indices,
token_to_seq_idx,
) = self._get_pruned_states(
hidden_states,
hidden_states_before_norm,
aux_hidden_states,
logits_metadata,
)
hidden_states_to_store = self._get_hidden_states_to_store(
hidden_states,
hidden_states_before_norm,
aux_hidden_states,
pruned_states,
pruned_states_before_norm,
aux_pruned_states,
sample_indices,
logits_metadata,
)
del hidden_states
if not logits_metadata.extend_return_logprob:
# Compute logits for both input and sampled tokens.
logits = self._get_logits(pruned_states, lm_head, logits_metadata)
sampled_logits = (
logits[sample_indices] if sample_indices is not None else logits
)
# Decode mode or extend mode without return_logprob.
return LogitsProcessorOutput(
next_token_logits=sampled_logits,
hidden_states=hidden_states_to_store,
)
# Start to process input logprobs
# Normalize the logprob w/o temperature, top-p
self._expand_metadata_for_logprobs(logits_metadata, pruned_states.device)
# Determine whether to use chunked or non-chunked logits processing.
# Skip chunking if:
# 1. Chunking is disabled
# 2. Total count is below chunk size threshold
# 3. DP attention all-gather is enabled (can use "enable_dp_lm_head" to enable chunking)
should_skip_chunking = (
not self.enable_logprobs_chunk
or pruned_states.shape[0] <= self.logprobs_chunk_size
or self.do_tensor_parallel_all_gather_dp_attn
)
if should_skip_chunking:
# Compute logits for both input and sampled tokens.
logits = self._get_logits(pruned_states, lm_head, logits_metadata)
sampled_logits = (
logits[sample_indices] if sample_indices is not None else logits
)
input_logits = logits[input_logprob_indices]
del logits
logprobs_result = self.process_input_logprobs(input_logits, logits_metadata)
else:
(logprobs_result, sampled_logits) = self.process_input_logprobs_by_chunk(
pruned_states,
sample_indices,
input_logprob_indices,
token_to_seq_idx,
lm_head,
logits_metadata,
)
return LogitsProcessorOutput(
next_token_logits=sampled_logits,
hidden_states=hidden_states_to_store,
input_token_logprobs=logprobs_result.input_token_logprobs,
input_top_logprobs_val=logprobs_result.input_top_logprobs_val,
input_top_logprobs_idx=logprobs_result.input_top_logprobs_idx,
input_token_ids_logprobs_val=logprobs_result.input_token_ids_logprobs_val,
input_token_ids_logprobs_idx=logprobs_result.input_token_ids_logprobs_idx,
)
def _get_pruned_states(
self,
hidden_states: torch.Tensor,
hidden_states_before_norm: Optional[torch.Tensor],
aux_hidden_states: Optional[torch.Tensor],
logits_metadata: LogitsMetadata,
):
pruned_states_before_norm: Optional[torch.Tensor] = None
aux_pruned_states = None
token_to_seq_idx = []
if (
logits_metadata.forward_mode.is_decode_or_idle()
or logits_metadata.forward_mode.is_target_verify()
@@ -473,7 +467,6 @@ class LogitsProcessor(nn.Module):
input_logprob_indices_pt = 0
input_logprob_indices = []
pt, pruned_states_list, pruned_states_before_norm_list = 0, [], []
token_to_seq_idx = []
for idx, (extend_logprob_start_len, extend_len) in enumerate(
zip(
@@ -524,12 +517,26 @@ class LogitsProcessor(nn.Module):
input_logprob_indices, device=pruned_states.device, dtype=torch.int64
)
full_logits = (
self._get_logits(hidden_states, lm_head, logits_metadata)
if self.return_full_logits
else None
return (
pruned_states,
pruned_states_before_norm,
aux_pruned_states,
sample_indices,
input_logprob_indices,
token_to_seq_idx,
)
def _get_hidden_states_to_store(
self,
hidden_states: torch.Tensor,
hidden_states_before_norm: Optional[torch.Tensor],
aux_hidden_states: Optional[List[torch.Tensor]],
pruned_states: torch.Tensor,
pruned_states_before_norm: Optional[torch.Tensor],
aux_pruned_states: Optional[List[torch.Tensor]],
sample_indices: Optional[torch.Tensor],
logits_metadata: LogitsMetadata,
) -> Optional[torch.Tensor]:
hidden_states_to_store: Optional[torch.Tensor] = None
hidden_states_to_store_before_norm: Optional[torch.Tensor] = None
if logits_metadata.capture_hidden_mode.need_capture():
@@ -565,32 +572,19 @@ class LogitsProcessor(nn.Module):
else:
assert False, "Should never reach"
del hidden_states
if hidden_states_to_store_before_norm is not None:
# NOTE: when hidden_states_before_norm is provided, we always
# prefer to return it.
hidden_states_to_store = hidden_states_to_store_before_norm
if not logits_metadata.extend_return_logprob:
# Compute logits for both input and sampled tokens.
logits = self._get_logits(pruned_states, lm_head, logits_metadata)
sampled_logits = (
logits[sample_indices] if sample_indices is not None else logits
)
return hidden_states_to_store
# Decode mode or extend mode without return_logprob.
return LogitsProcessorOutput(
full_logits=full_logits,
next_token_logits=sampled_logits,
hidden_states=hidden_states_to_store,
)
# Start to process input logprobs
# Normalize the logprob w/o temperature, top-p
def _expand_metadata_for_logprobs(
self, logits_metadata: LogitsMetadata, device: torch.device
):
pruned_lens = torch.tensor(
logits_metadata.extend_logprob_pruned_lens_cpu,
device=pruned_states.device,
device=device,
)
if logits_metadata.temp_scaled_logprobs:
logits_metadata.temperature = torch.repeat_interleave(
@@ -603,49 +597,6 @@ class LogitsProcessor(nn.Module):
pruned_lens,
)
# Determine whether to use chunked or non-chunked logits processing.
# Skip chunking if:
# 1. Chunking is disabled
# 2. Total count is below chunk size threshold
# 3. DP attention all-gather is enabled (can use "enable_dp_lm_head" to enable chunking)
should_skip_chunking = (
not self.enable_logprobs_chunk
or pruned_states.shape[0] <= self.logprobs_chunk_size
or self.do_tensor_parallel_all_gather_dp_attn
)
if should_skip_chunking:
# Compute logits for both input and sampled tokens.
logits = self._get_logits(pruned_states, lm_head, logits_metadata)
sampled_logits = (
logits[sample_indices] if sample_indices is not None else logits
)
input_logits = logits[input_logprob_indices]
del logits
logprobs_result = self.process_input_logprobs(input_logits, logits_metadata)
else:
(logprobs_result, sampled_logits) = self.process_input_logprobs_by_chunk(
pruned_states,
sample_indices,
input_logprob_indices,
token_to_seq_idx,
lm_head,
logits_metadata,
)
return LogitsProcessorOutput(
full_logits=full_logits,
next_token_logits=sampled_logits,
hidden_states=hidden_states_to_store,
input_token_logprobs=logprobs_result.input_token_logprobs,
input_top_logprobs_val=logprobs_result.input_top_logprobs_val,
input_top_logprobs_idx=logprobs_result.input_top_logprobs_idx,
input_token_ids_logprobs_val=logprobs_result.input_token_ids_logprobs_val,
input_token_ids_logprobs_idx=logprobs_result.input_token_ids_logprobs_idx,
)
def process_input_logprobs(self, input_logits, logits_metadata: LogitsMetadata):
input_logprobs = compute_temp_top_p_normalized_logprobs(
input_logits, logits_metadata
@@ -860,18 +811,48 @@ class LogitsProcessor(nn.Module):
last position (e.g., extend without input logprobs). The caller should
guarantee the given hidden_states follow this constraint.
"""
if self.do_tensor_parallel_all_gather_dp_attn:
logits_metadata.compute_dp_attention_metadata()
hidden_states, local_hidden_states = (
logits_metadata.gathered_buffer,
hidden_states,
)
dp_gather_replicate(hidden_states, local_hidden_states, logits_metadata)
hidden_states, local_hidden_states = self._gather_dp_attn_hidden_states(
hidden_states, logits_metadata
)
logits = self._compute_lm_head(hidden_states, lm_head, embedding_bias)
if self.logit_scale is not None:
logits.mul_(self.logit_scale)
if self.do_tensor_parallel_all_gather:
if self.use_attn_tp_group:
logits = self._gather_attn_tp_logits(logits)
else:
logits = tensor_model_parallel_all_gather(logits)
logits = self._scatter_dp_attn_logits(
logits, local_hidden_states, logits_metadata
)
logits = self._copy_logits_to_buffer(logits, logits_metadata)
if self.final_logit_softcapping:
if not _is_npu:
fused_softcap(logits, self.final_logit_softcapping)
else:
logits = self.final_logit_softcapping * torch.tanh(
logits / self.final_logit_softcapping
)
return logits
def _compute_lm_head(
self,
hidden_states: torch.Tensor,
lm_head: VocabParallelEmbedding,
embedding_bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
if hasattr(lm_head, "set_lora") and hasattr(lm_head, "apply_lora"):
# This is a LoRA-wrapped module, use its forward method
logits = lm_head(hidden_states)
elif hasattr(lm_head, "weight"):
# Normal linear layer
if self.use_fp32_lm_head:
logits = torch.matmul(
hidden_states.to(torch.float32), lm_head.weight.to(torch.float32).T
@@ -904,52 +885,66 @@ class LogitsProcessor(nn.Module):
logits = lm_head.quant_method.apply(
lm_head, hidden_states, embedding_bias
)
return logits
if self.logit_scale is not None:
logits.mul_(self.logit_scale)
if self.do_tensor_parallel_all_gather:
if self.use_attn_tp_group:
if self.config.vocab_size % self.attn_tp_size == 0:
global_logits = torch.empty(
(
self.attn_tp_size,
logits.shape[0],
self.config.vocab_size // self.attn_tp_size,
),
device=logits.device,
dtype=logits.dtype,
)
attn_tp_all_gather_into_tensor(global_logits, logits)
global_logits = global_logits.permute(1, 0, 2).reshape(
logits.shape[0], self.config.vocab_size
)
else:
global_logits = torch.empty(
(self.config.vocab_size, logits.shape[0]),
device=logits.device,
dtype=logits.dtype,
)
global_logits = global_logits.T
attn_tp_all_gather(
list(global_logits.tensor_split(self.attn_tp_size, dim=-1)),
logits,
)
logits = global_logits
else:
logits = tensor_model_parallel_all_gather(logits)
def _gather_dp_attn_hidden_states(
self, hidden_states: torch.Tensor, logits_metadata: LogitsMetadata
) -> Tuple[torch.Tensor, torch.Tensor]:
if self.do_tensor_parallel_all_gather_dp_attn:
logits, global_logits = (
torch.empty(
(local_hidden_states.shape[0], logits.shape[1]),
device=logits.device,
dtype=logits.dtype,
logits_metadata.compute_dp_attention_metadata()
local_hidden_states = hidden_states
hidden_states = logits_metadata.gathered_buffer
dp_gather_replicate(hidden_states, local_hidden_states, logits_metadata)
return hidden_states, local_hidden_states
return hidden_states, hidden_states
def _gather_attn_tp_logits(self, logits: torch.Tensor) -> torch.Tensor:
if self.config.vocab_size % self.attn_tp_size == 0:
global_logits = torch.empty(
(
self.attn_tp_size,
logits.shape[0],
self.config.vocab_size // self.attn_tp_size,
),
device=logits.device,
dtype=logits.dtype,
)
attn_tp_all_gather_into_tensor(global_logits, logits)
global_logits = global_logits.permute(1, 0, 2).reshape(
logits.shape[0], self.config.vocab_size
)
else:
global_logits = torch.empty(
(self.config.vocab_size, logits.shape[0]),
device=logits.device,
dtype=logits.dtype,
)
global_logits = global_logits.T
attn_tp_all_gather(
list(global_logits.tensor_split(self.attn_tp_size, dim=-1)),
logits,
)
dp_scatter(logits, global_logits, logits_metadata)
return global_logits
def _scatter_dp_attn_logits(
self,
logits: torch.Tensor,
local_hidden_states: torch.Tensor,
logits_metadata: LogitsMetadata,
) -> torch.Tensor:
if self.do_tensor_parallel_all_gather_dp_attn:
global_logits = logits
logits = torch.empty(
(local_hidden_states.shape[0], global_logits.shape[1]),
device=global_logits.device,
dtype=global_logits.dtype,
)
dp_scatter(logits, global_logits, logits_metadata)
return logits
def _copy_logits_to_buffer(
self, logits: torch.Tensor, logits_metadata: LogitsMetadata
) -> torch.Tensor:
if logits_metadata.next_token_logits_buffer is not None:
logits_buffer = logits_metadata.next_token_logits_buffer
assert logits_buffer.dtype == torch.float
@@ -957,17 +952,123 @@ class LogitsProcessor(nn.Module):
logits = logits_buffer
else:
logits = logits[:, : self.config.vocab_size].float()
if self.final_logit_softcapping:
if not _is_npu:
fused_softcap(logits, self.final_logit_softcapping)
else:
logits = self.final_logit_softcapping * torch.tanh(
logits / self.final_logit_softcapping
)
return logits
def _get_dllm_logits(
self,
hidden_states: torch.Tensor,
lm_head: VocabParallelEmbedding,
logits_metadata: LogitsMetadata,
) -> LogitsProcessorOutput:
assert self.return_full_logits
full_logits = self._get_logits(hidden_states, lm_head, logits_metadata)
return LogitsProcessorOutput(
full_logits=full_logits,
next_token_logits=None,
)
def compute_logprobs_for_multi_item_scoring(
self,
input_ids,
hidden_states,
lm_head: VocabParallelEmbedding,
logits_metadata: Union[LogitsMetadata, ForwardBatch],
delimiter_token: int,
):
"""
Compute logprobs for multi-item scoring using delimiter-based token extraction.
This method is designed for scenarios where you want to score multiple items/candidates
against a single query by combining them into one sequence separated by delimiters.
Sequence format: Query<delimiter>Item1<delimiter>Item2<delimiter>...
Scoring positions: Extracts logprobs at positions before each <delimiter>
Args:
input_ids (torch.Tensor): Input token IDs containing query and items separated by delimiters.
Shape: [total_sequence_length] for single request or [batch_total_length] for batch.
hidden_states (torch.Tensor): Hidden states from the model.
Shape: [sequence_length, hidden_dim].
lm_head (VocabParallelEmbedding): Language model head for computing logits.
logits_metadata (Union[LogitsMetadata, ForwardBatch]): Metadata containing batch info
and token ID specifications for logprob extraction.
delimiter_token (int): Token ID used as delimiter between query and items.
Returns:
LogitsProcessorOutput: Contains:
- next_token_logits: None (not needed for scoring-only requests)
- input_token_logprobs: Logprobs of delimiter tokens at scoring positions
- input_top_logprobs_val: Top-k logprobs at delimiter positions (if requested)
- input_top_logprobs_idx: Top-k token indices at delimiter positions (if requested)
- input_token_ids_logprobs_val: Logprobs for user-requested token IDs (if any)
- input_token_ids_logprobs_idx: Indices for user-requested token IDs (if any)
"""
multi_item_indices = (input_ids == delimiter_token).nonzero(as_tuple=True)[
0
] - 1
# Extract hidden states at delimiter positions for multi-item scoring
sliced_hidden = hidden_states[multi_item_indices]
sliced_logits = self._get_logits(sliced_hidden, lm_head, logits_metadata)
sliced_logprobs = torch.nn.functional.log_softmax(sliced_logits, dim=-1)
# Initialize return values
input_token_ids_logprobs_val = []
input_token_ids_logprobs_idx = []
input_top_logprobs_val = None
input_top_logprobs_idx = None
# Recalculate extend_logprob_pruned_lens_cpu to match delimiter counts per request
# Original contains sequence lengths, but we need delimiter counts for sliced_logprobs
if (
logits_metadata.token_ids_logprobs
or logits_metadata.extend_return_top_logprob
):
logits_metadata.extend_logprob_pruned_lens_cpu = []
if logits_metadata.extend_seq_lens_cpu is not None:
# Multi-request batch: count delimiters per request
input_pt = 0
for req_seq_len in logits_metadata.extend_seq_lens_cpu:
req_input_ids = input_ids[input_pt : input_pt + req_seq_len]
delimiter_count = (req_input_ids == delimiter_token).sum().item()
logits_metadata.extend_logprob_pruned_lens_cpu.append(
delimiter_count
)
input_pt += req_seq_len
else:
# Single request case: one request gets all delimiters
total_delimiters = (input_ids == delimiter_token).sum().item()
logits_metadata.extend_logprob_pruned_lens_cpu = [total_delimiters]
# Get the logprobs of specified token ids
if logits_metadata.extend_token_ids_logprob:
(
input_token_ids_logprobs_val,
input_token_ids_logprobs_idx,
) = get_token_ids_logprobs_prefill(
sliced_logprobs, logits_metadata, delay_cpu_copy=True
)
# Get the logprob of top-k tokens
if logits_metadata.extend_return_top_logprob:
(
input_top_logprobs_val,
input_top_logprobs_idx,
) = get_top_logprobs_prefill(sliced_logprobs, logits_metadata)
# For input_token_logprobs, use delimiter token logprobs
input_token_logprobs = sliced_logprobs[:, delimiter_token]
return LogitsProcessorOutput(
next_token_logits=None, # Multi-item scoring doesn't need next token logits
input_token_logprobs=input_token_logprobs,
input_top_logprobs_val=input_top_logprobs_val,
input_top_logprobs_idx=input_top_logprobs_idx,
input_token_ids_logprobs_val=input_token_ids_logprobs_val,
input_token_ids_logprobs_idx=input_token_ids_logprobs_idx,
)
@triton.jit
def fused_softcap_kernel(

View File

@@ -126,52 +126,9 @@ class Sampler(nn.Module):
probs = logits
del logits
if can_sample_directly_from_probs:
# when we don't need top-k, top-p, or min-p sampling, we can directly sample from the probs
batch_next_token_ids = sampling_from_probs_torch(
probs,
sampling_seed=sampling_info.sampling_seed,
positions=positions,
)
else:
if get_global_server_args().sampling_backend == "flashinfer":
if sampling_info.need_min_p_sampling:
probs = top_k_renorm_prob(probs, sampling_info.top_ks)
probs = top_p_renorm_prob(probs, sampling_info.top_ps)
batch_next_token_ids = min_p_sampling_from_probs(
probs, sampling_info.min_ps
)
else:
batch_next_token_ids = top_k_top_p_sampling_from_probs(
probs.contiguous(),
sampling_info.top_ks,
sampling_info.top_ps,
filter_apply_order="joint",
check_nan=self.use_nan_detection,
)
elif get_global_server_args().sampling_backend == "pytorch":
# A slower fallback implementation with torch native operations.
batch_next_token_ids = top_k_top_p_min_p_sampling_from_probs_torch(
probs,
sampling_info.top_ks,
sampling_info.top_ps,
sampling_info.min_ps,
sampling_info.need_min_p_sampling,
sampling_info.sampling_seed,
positions,
)
elif get_global_server_args().sampling_backend == "ascend":
batch_next_token_ids = top_k_top_p_min_p_sampling_from_probs_ascend(
probs,
sampling_info.top_ks,
sampling_info.top_ps,
sampling_info.min_ps,
sampling_info.need_min_p_sampling,
)
else:
raise ValueError(
f"Invalid sampling backend: {get_global_server_args().sampling_backend}"
)
batch_next_token_ids = self._sample_from_probs(
probs, sampling_info, positions, can_sample_directly_from_probs
)
if return_logprob:
if get_global_server_args().rl_on_policy_target is not None:
@@ -188,23 +145,77 @@ class Sampler(nn.Module):
# Attach logprobs to logits_output (in-place modification)
if return_logprob:
if any(x > 0 for x in top_logprobs_nums):
(
logits_output.next_token_top_logprobs_val,
logits_output.next_token_top_logprobs_idx,
) = get_top_logprobs(logprobs, top_logprobs_nums)
if any(x is not None for x in token_ids_logprobs):
(
logits_output.next_token_token_ids_logprobs_val,
logits_output.next_token_token_ids_logprobs_idx,
) = get_token_ids_logprobs(logprobs, token_ids_logprobs)
logits_output.next_token_logprobs = logprobs[
torch.arange(len(batch_next_token_ids), device=sampling_info.device),
self._attach_logprobs_to_output(
logits_output,
logprobs,
top_logprobs_nums,
token_ids_logprobs,
sampling_info,
batch_next_token_ids,
]
)
self._sync_token_ids_across_tp(batch_next_token_ids, sampling_info)
return batch_next_token_ids
def _sample_from_probs(
self,
probs: torch.Tensor,
sampling_info: SamplingBatchInfo,
positions: torch.Tensor,
can_sample_directly_from_probs: bool,
) -> torch.Tensor:
if can_sample_directly_from_probs:
# when we don't need top-k, top-p, or min-p sampling, we can directly sample from the probs
batch_next_token_ids = sampling_from_probs_torch(
probs,
sampling_seed=sampling_info.sampling_seed,
positions=positions,
)
else:
if get_global_server_args().sampling_backend == "flashinfer":
if sampling_info.need_min_p_sampling:
probs = top_k_renorm_prob(probs, sampling_info.top_ks)
probs = top_p_renorm_prob(probs, sampling_info.top_ps)
batch_next_token_ids = min_p_sampling_from_probs(
probs, sampling_info.min_ps
)
else:
batch_next_token_ids = top_k_top_p_sampling_from_probs(
probs.contiguous(),
sampling_info.top_ks,
sampling_info.top_ps,
filter_apply_order="joint",
check_nan=self.use_nan_detection,
)
elif get_global_server_args().sampling_backend == "pytorch":
# A slower fallback implementation with torch native operations.
batch_next_token_ids = top_k_top_p_min_p_sampling_from_probs_torch(
probs,
sampling_info.top_ks,
sampling_info.top_ps,
sampling_info.min_ps,
sampling_info.need_min_p_sampling,
sampling_info.sampling_seed,
positions,
)
elif get_global_server_args().sampling_backend == "ascend":
batch_next_token_ids = top_k_top_p_min_p_sampling_from_probs_ascend(
probs,
sampling_info.top_ks,
sampling_info.top_ps,
sampling_info.min_ps,
sampling_info.need_min_p_sampling,
)
else:
raise ValueError(
f"Invalid sampling backend: {get_global_server_args().sampling_backend}"
)
return batch_next_token_ids
def _sync_token_ids_across_tp(
self, batch_next_token_ids: torch.Tensor, sampling_info: SamplingBatchInfo
):
if SYNC_TOKEN_IDS_ACROSS_TP or sampling_info.grammars:
# For performance reasons, SGLang does not sync the final token IDs across TP ranks by default.
# This saves one all-reduce, but the correctness of this approach depends on the determinism of several operators:
@@ -219,7 +230,32 @@ class Sampler(nn.Module):
group=self.tp_sync_group,
)
return batch_next_token_ids
def _attach_logprobs_to_output(
self,
logits_output: LogitsProcessorOutput,
logprobs: torch.Tensor,
top_logprobs_nums: List[int],
token_ids_logprobs: List[List[int]],
sampling_info: SamplingBatchInfo,
batch_next_token_ids: torch.Tensor,
):
# Attach logprobs to logits_output (in-place modification)
if any(x > 0 for x in top_logprobs_nums):
(
logits_output.next_token_top_logprobs_val,
logits_output.next_token_top_logprobs_idx,
) = get_top_logprobs(logprobs, top_logprobs_nums)
if any(x is not None for x in token_ids_logprobs):
(
logits_output.next_token_token_ids_logprobs_val,
logits_output.next_token_token_ids_logprobs_idx,
) = get_token_ids_logprobs(logprobs, token_ids_logprobs)
logits_output.next_token_logprobs = logprobs[
torch.arange(len(batch_next_token_ids), device=sampling_info.device),
batch_next_token_ids,
]
def compute_logprobs_only(
self,