[Generative Score API] Scoring(Prefill-only) optimizations. (#9748)
This commit is contained in:
@@ -1,5 +1,5 @@
|
||||
import logging
|
||||
from typing import List
|
||||
from typing import List, Tuple
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
@@ -39,6 +39,25 @@ class Sampler(nn.Module):
|
||||
if is_dp_attention_enabled():
|
||||
self.tp_sync_group = get_attention_tp_group().device_group
|
||||
|
||||
def _preprocess_logits(
|
||||
self, logits: torch.Tensor, sampling_info: SamplingBatchInfo
|
||||
) -> torch.Tensor:
|
||||
"""Apply custom logit processors and handle NaN detection."""
|
||||
# Apply the custom logit processors if registered in the sampling info
|
||||
if sampling_info.has_custom_logit_processor:
|
||||
apply_custom_logit_processor(logits, sampling_info)
|
||||
|
||||
# Detect and handle NaN values in logits
|
||||
if self.use_nan_detection and torch.any(torch.isnan(logits)):
|
||||
logger.warning("Detected errors during sampling! NaN in the logits.")
|
||||
logits = torch.where(
|
||||
torch.isnan(logits), torch.full_like(logits, -1e5), logits
|
||||
)
|
||||
if crash_on_warnings():
|
||||
raise ValueError("Detected errors during sampling! NaN in the logits.")
|
||||
|
||||
return logits
|
||||
|
||||
def forward(
|
||||
self,
|
||||
logits_output: LogitsProcessorOutput,
|
||||
@@ -61,17 +80,8 @@ class Sampler(nn.Module):
|
||||
"""
|
||||
logits = logits_output.next_token_logits
|
||||
|
||||
# Apply the custom logit processors if registered in the sampling info.
|
||||
if sampling_info.has_custom_logit_processor:
|
||||
apply_custom_logit_processor(logits, sampling_info)
|
||||
|
||||
if self.use_nan_detection and torch.any(torch.isnan(logits)):
|
||||
logger.warning("Detected errors during sampling! NaN in the logits.")
|
||||
logits = torch.where(
|
||||
torch.isnan(logits), torch.full_like(logits, -1e5), logits
|
||||
)
|
||||
if crash_on_warnings():
|
||||
raise ValueError("Detected errors during sampling! NaN in the logits.")
|
||||
# Preprocess logits (custom processors and NaN handling)
|
||||
logits = self._preprocess_logits(logits, sampling_info)
|
||||
|
||||
if sampling_info.is_all_greedy:
|
||||
# Use torch.argmax if all requests use greedy sampling
|
||||
@@ -165,6 +175,54 @@ class Sampler(nn.Module):
|
||||
|
||||
return batch_next_token_ids
|
||||
|
||||
def compute_logprobs_only(
|
||||
self,
|
||||
logits_output: LogitsProcessorOutput,
|
||||
sampling_info: SamplingBatchInfo,
|
||||
return_logprob: bool,
|
||||
top_logprobs_nums: List[int],
|
||||
token_ids_logprobs: List[List[int]],
|
||||
) -> None:
|
||||
"""
|
||||
Compute logprobs for requested token IDs without performing sampling.
|
||||
|
||||
Optimized for prefill-only scoring requests that need token probabilities
|
||||
but don't require next token generation.
|
||||
"""
|
||||
if logits_output.next_token_logits is None:
|
||||
logger.warning("No logits available for logprob computation")
|
||||
return
|
||||
|
||||
# Check if any requests actually need logprobs computation
|
||||
needs_token_ids_logprobs = any(
|
||||
token_ids is not None and len(token_ids) > 0
|
||||
for token_ids in token_ids_logprobs
|
||||
)
|
||||
needs_top_logprobs = any(x > 0 for x in top_logprobs_nums)
|
||||
|
||||
if not (needs_token_ids_logprobs or needs_top_logprobs):
|
||||
return
|
||||
|
||||
# Preprocess logits (custom processors and NaN handling)
|
||||
logits = self._preprocess_logits(logits_output.next_token_logits, sampling_info)
|
||||
|
||||
# Compute logprobs
|
||||
logprobs = torch.nn.functional.log_softmax(logits, dim=-1)
|
||||
|
||||
# Handle top logprobs if requested
|
||||
if needs_top_logprobs:
|
||||
(
|
||||
logits_output.next_token_top_logprobs_val,
|
||||
logits_output.next_token_top_logprobs_idx,
|
||||
) = get_top_logprobs(logprobs, top_logprobs_nums)
|
||||
|
||||
# Handle token_ids logprobs if requested
|
||||
if needs_token_ids_logprobs:
|
||||
(
|
||||
logits_output.next_token_token_ids_logprobs_val,
|
||||
logits_output.next_token_token_ids_logprobs_idx,
|
||||
) = get_token_ids_logprobs_batch_optimized(logprobs, token_ids_logprobs)
|
||||
|
||||
|
||||
def top_k_top_p_min_p_sampling_from_probs_torch(
|
||||
probs: torch.Tensor,
|
||||
@@ -234,10 +292,95 @@ def get_top_logprobs(
|
||||
)
|
||||
|
||||
|
||||
def get_token_ids_logprobs(
|
||||
def get_token_ids_logprobs_batch_optimized(
|
||||
logprobs: torch.Tensor,
|
||||
token_ids_logprobs: List[List[int]],
|
||||
):
|
||||
) -> Tuple[List, List]:
|
||||
"""
|
||||
Vectorized batch processing for token ID logprobs extraction.
|
||||
|
||||
Uses a single GPU kernel call for the entire batch instead of multiple
|
||||
separate calls, significantly improving performance for large batches.
|
||||
|
||||
Args:
|
||||
logprobs: Log probabilities tensor [batch_size, vocab_size]
|
||||
token_ids_logprobs: List of token IDs to extract logprobs for
|
||||
|
||||
Example:
|
||||
# Input: batch_size=3, vocab_size=5
|
||||
logprobs = torch.tensor([
|
||||
[-1.2, -2.1, -0.8, -3.0, -1.5], # batch 0
|
||||
[-0.5, -1.8, -2.2, -1.1, -2.7], # batch 1
|
||||
[-2.0, -0.9, -1.4, -2.8, -1.6], # batch 2
|
||||
])
|
||||
token_ids_logprobs = [[1, 3], [2], [0, 2, 4]]
|
||||
|
||||
# Output:
|
||||
# values = [tensor([-2.1, -3.0]), tensor([-2.2]), tensor([-2.0, -1.4, -1.6])]
|
||||
# indices = [[1, 3], [2], [0, 2, 4]]
|
||||
"""
|
||||
batch_size = len(token_ids_logprobs)
|
||||
device = logprobs.device
|
||||
|
||||
# Step 1: Calculate lengths for each request, treating None as empty list
|
||||
# Example: [[1, 3], [2], [0, 2, 4]] -> token_lengths = tensor([2, 1, 3])
|
||||
token_lengths = torch.tensor(
|
||||
[len(token_ids or []) for token_ids in token_ids_logprobs], device=device
|
||||
)
|
||||
total_tokens = int(token_lengths.sum().item()) # 2 + 1 + 3 = 6
|
||||
|
||||
# Handle edge case where no tokens are requested
|
||||
if total_tokens == 0:
|
||||
return [logprobs.new_empty(0) for _ in token_ids_logprobs], [
|
||||
[] for _ in token_ids_logprobs
|
||||
]
|
||||
|
||||
# Step 2: Build flattened indices using torch operations
|
||||
# Example: row_indices = [0, 0, 1, 2, 2, 2] (batch indices repeated by their lengths)
|
||||
row_indices = torch.repeat_interleave(
|
||||
torch.arange(batch_size, device=device), token_lengths
|
||||
)
|
||||
# Example: col_indices = [1, 3, 2, 0, 2, 4] (flattened token IDs from all requests)
|
||||
col_indices = torch.tensor(
|
||||
[
|
||||
token_id
|
||||
for token_ids in token_ids_logprobs
|
||||
for token_id in (token_ids or [])
|
||||
],
|
||||
device=device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
|
||||
# Step 3: Single vectorized gather operation
|
||||
# Example: logprobs[row_indices, col_indices] -> [-2.1, -3.0, -2.2, -2.0, -1.4, -1.6]
|
||||
gathered_logprobs = logprobs[row_indices, col_indices]
|
||||
|
||||
# Step 4: Split results back per request using torch operations
|
||||
# Example: split tensor [6] into chunks of sizes [2, 1, 3] -> [tensor(2), tensor(1), tensor(3)]
|
||||
split_logprobs = torch.split_with_sizes(
|
||||
gathered_logprobs, token_lengths.tolist(), dim=0
|
||||
)
|
||||
|
||||
# Step 5: Format output to match expected return structure
|
||||
# Example: Convert split tensors back to list format with proper empty handling
|
||||
# i=0: [1,3] -> append split_logprobs[0] and [1,3]
|
||||
# i=1: [2] -> append split_logprobs[1] and [2]
|
||||
# i=2: [0,2,4] -> append split_logprobs[2] and [0,2,4]
|
||||
output_token_ids_logprobs_val = []
|
||||
output_token_ids_logprobs_idx = []
|
||||
|
||||
for i, token_ids in enumerate(token_ids_logprobs):
|
||||
if token_ids is not None and len(token_ids) > 0:
|
||||
output_token_ids_logprobs_val.append(split_logprobs[i])
|
||||
output_token_ids_logprobs_idx.append(token_ids)
|
||||
else:
|
||||
output_token_ids_logprobs_val.append(logprobs.new_empty(0))
|
||||
output_token_ids_logprobs_idx.append([])
|
||||
|
||||
return output_token_ids_logprobs_val, output_token_ids_logprobs_idx
|
||||
|
||||
|
||||
def get_token_ids_logprobs(logprobs: torch.Tensor, token_ids_logprobs: List[List[int]]):
|
||||
output_token_ids_logprobs_val = []
|
||||
output_token_ids_logprobs_idx = []
|
||||
for i, token_ids in enumerate(token_ids_logprobs):
|
||||
|
||||
Reference in New Issue
Block a user