Make the functions in logits_processor.py and sampler.py more modular (#17885)
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user