[1 / N] Clean up logprob utils (#15509)

This commit is contained in:
Liangsheng Yin
2025-12-22 03:12:25 +08:00
committed by GitHub
parent 1d9ba2ce4d
commit 8766a1dd24
8 changed files with 334 additions and 311 deletions

View File

@@ -40,6 +40,14 @@ from sglang.srt.layers.dp_attention import (
get_dp_dtype,
get_dp_hidden_size,
)
from sglang.srt.layers.utils.logprob import (
InputLogprobsResult,
compute_temp_top_p_normalized_logprobs,
get_token_ids_logprobs_chunk,
get_token_ids_logprobs_prefill,
get_top_logprobs_chunk,
get_top_logprobs_prefill,
)
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode,
@@ -54,15 +62,6 @@ logger = logging.getLogger(__name__)
_is_npu = is_npu()
@dataclasses.dataclass
class InputLogprobsResult:
input_token_logprobs: torch.Tensor
input_top_logprobs_val: Optional[List] = None
input_top_logprobs_idx: Optional[List] = None
input_token_ids_logprobs_val: Optional[List] = None
input_token_ids_logprobs_idx: Optional[List] = None
@dataclasses.dataclass
class LogitsProcessorOutput:
## Part 1: This part will be assigned in python/sglang/srt/layers/logits_processor.py::LogitsProcessor
@@ -351,7 +350,7 @@ class LogitsProcessor(nn.Module):
(
input_token_ids_logprobs_val,
input_token_ids_logprobs_idx,
) = self.get_token_ids_logprobs(
) = get_token_ids_logprobs_prefill(
sliced_logprobs, logits_metadata, delay_cpu_copy=True
)
@@ -360,7 +359,7 @@ class LogitsProcessor(nn.Module):
(
input_top_logprobs_val,
input_top_logprobs_idx,
) = self.get_top_logprobs(sliced_logprobs, logits_metadata)
) = get_top_logprobs_prefill(sliced_logprobs, logits_metadata)
# For input_token_logprobs, use delimiter token logprobs
input_token_logprobs = sliced_logprobs[:, delimiter_token]
@@ -619,14 +618,12 @@ class LogitsProcessor(nn.Module):
logits[sample_indices] if sample_indices is not None else logits
)
input_logprobs = logits[input_logprob_indices]
input_logits = logits[input_logprob_indices]
del logits
logprobs_result = self._process_input_logprobs(
input_logprobs, logits_metadata
)
logprobs_result = self.process_input_logprobs(input_logits, logits_metadata)
else:
(logprobs_result, sampled_logits) = self._process_input_logprobs_by_chunk(
(logprobs_result, sampled_logits) = self.process_input_logprobs_by_chunk(
pruned_states,
sample_indices,
input_logprob_indices,
@@ -646,9 +643,9 @@ class LogitsProcessor(nn.Module):
input_token_ids_logprobs_idx=logprobs_result.input_token_ids_logprobs_idx,
)
def _process_input_logprobs(self, input_logprobs, logits_metadata):
input_logprobs = self.compute_temp_top_p_normalized_logprobs(
input_logprobs, logits_metadata
def process_input_logprobs(self, input_logits, logits_metadata: LogitsMetadata):
input_logprobs = compute_temp_top_p_normalized_logprobs(
input_logits, logits_metadata
)
# Get the logprob of top-k tokens
@@ -656,7 +653,7 @@ class LogitsProcessor(nn.Module):
(
input_top_logprobs_val,
input_top_logprobs_idx,
) = self.get_top_logprobs(input_logprobs, logits_metadata)
) = get_top_logprobs_prefill(input_logprobs, logits_metadata)
else:
input_top_logprobs_val = input_top_logprobs_idx = None
@@ -665,7 +662,7 @@ class LogitsProcessor(nn.Module):
(
input_token_ids_logprobs_val,
input_token_ids_logprobs_idx,
) = self.get_token_ids_logprobs(input_logprobs, logits_metadata)
) = get_token_ids_logprobs_prefill(input_logprobs, logits_metadata)
else:
input_token_ids_logprobs_val = input_token_ids_logprobs_idx = None
@@ -682,7 +679,7 @@ class LogitsProcessor(nn.Module):
input_token_ids_logprobs_idx=input_token_ids_logprobs_idx,
)
def _process_input_logprobs_by_chunk(
def process_input_logprobs_by_chunk(
self,
pruned_states: torch.Tensor,
sample_indices: torch.Tensor,
@@ -776,7 +773,7 @@ class LogitsProcessor(nn.Module):
if logits_metadata.top_p is not None
else None
)
chunk_input_logprobs = self.compute_temp_top_p_normalized_logprobs(
chunk_input_logprobs = compute_temp_top_p_normalized_logprobs(
chunk_input_logprobs,
logits_metadata,
chunk_top_p,
@@ -794,7 +791,7 @@ class LogitsProcessor(nn.Module):
pruned_lens = logits_metadata.extend_logprob_pruned_lens_cpu[
chunk_slice
]
split_len_topk = self.get_top_logprobs_chunk(
split_len_topk = get_top_logprobs_chunk(
chunk_input_logprobs,
logits_metadata,
top_k_nums,
@@ -810,9 +807,8 @@ class LogitsProcessor(nn.Module):
pruned_lens = logits_metadata.extend_logprob_pruned_lens_cpu[
chunk_slice
]
split_len_token_ids = self.get_token_ids_logprobs_chunk(
split_len_token_ids = get_token_ids_logprobs_chunk(
chunk_input_logprobs,
logits_metadata,
token_ids_logprobs,
pruned_lens,
input_token_ids_logprobs_val,
@@ -964,250 +960,6 @@ class LogitsProcessor(nn.Module):
return logits
@staticmethod
def get_top_logprobs(all_logprobs: torch.Tensor, logits_metadata: LogitsMetadata):
max_k = max(logits_metadata.top_logprobs_nums)
ret = all_logprobs.topk(max_k, dim=1)
values = ret.values.tolist()
indices = ret.indices.tolist()
input_top_logprobs_val, input_top_logprobs_idx = [], []
pt = 0
for k, pruned_len in zip(
logits_metadata.top_logprobs_nums,
logits_metadata.extend_logprob_pruned_lens_cpu,
):
if pruned_len <= 0:
input_top_logprobs_val.append([])
input_top_logprobs_idx.append([])
continue
input_top_logprobs_val.append(
[values[pt + j][:k] for j in range(pruned_len)]
)
input_top_logprobs_idx.append(
[indices[pt + j][:k] for j in range(pruned_len)]
)
pt += pruned_len
return input_top_logprobs_val, input_top_logprobs_idx
@staticmethod
def get_top_logprobs_chunk(
logprobs: torch.Tensor,
logits_metadata: LogitsMetadata,
top_k_nums: List[int],
pruned_lens: List[int],
input_top_logprobs_val: List,
input_top_logprobs_idx: List,
split_pruned_len: int,
) -> int:
"""Get top-k logprobs for each sequence in the chunk.
Args:
logprobs: Log probabilities tensor of shape [seq_len, vocab_size]
logits_metadata: Metadata containing top-k and pruned length info
top_k_nums: List of top-k numbers for each sequence
pruned_lens: List of pruned lengths for each sequence
input_top_logprobs_val: List to store top-k logprob values
input_top_logprobs_idx: List to store top-k token indices
split_pruned_len: Length of pruned tokens from previous chunk
Returns:
int: Number of remaining tokens to process in next chunk
"""
# No sequences in the chunk
if logprobs.shape[0] == 0:
return 0
max_k = max(logits_metadata.top_logprobs_nums)
ret = logprobs.topk(max_k, dim=1)
values = ret.values.tolist()
indices = ret.indices.tolist()
pt = 0
next_split_pruned_len = 0
for n, (k, pruned_len) in enumerate(zip(top_k_nums, pruned_lens)):
if n == 0:
# For the first sequence, adjust the pruned length
pruned_len -= split_pruned_len
else:
# After the first sequence, no split in the middle
split_pruned_len = 0
if pruned_len <= 0:
# if pruned length is less than or equal to 0,
# there is no top-k logprobs to process
input_top_logprobs_val.append([])
input_top_logprobs_idx.append([])
continue
# Get the top-k logprobs
val = []
idx = []
for j in range(pruned_len):
# Handle remaining tokens in next chunk if any
if pt + j >= len(values):
next_split_pruned_len = split_pruned_len + j
break
# Append the top-k logprobs
val.append(values[pt + j][:k])
idx.append(indices[pt + j][:k])
# Append or extend based on whether the sequence was split across chunks
if len(val) > 0:
if split_pruned_len > 0:
input_top_logprobs_val[-1].extend(val)
input_top_logprobs_idx[-1].extend(idx)
else:
input_top_logprobs_val.append(val)
input_top_logprobs_idx.append(idx)
pt += pruned_len
return next_split_pruned_len
@staticmethod
def get_token_ids_logprobs(
all_logprobs: torch.Tensor,
logits_metadata: LogitsMetadata,
delay_cpu_copy: bool = False,
):
input_token_ids_logprobs_val, input_token_ids_logprobs_idx = [], []
pt = 0
for token_ids, pruned_len in zip(
logits_metadata.token_ids_logprobs,
logits_metadata.extend_logprob_pruned_lens_cpu,
):
if pruned_len <= 0:
input_token_ids_logprobs_val.append([])
input_token_ids_logprobs_idx.append([])
continue
position_logprobs = all_logprobs[
pt : pt + pruned_len, token_ids
] # Shape: [pruned_len, num_tokens]
if delay_cpu_copy:
# Keep as tensor to delay GPU-to-CPU transfer
input_token_ids_logprobs_val.append(position_logprobs)
else:
# Convert to list immediately (default behavior)
input_token_ids_logprobs_val.append(position_logprobs.tolist())
input_token_ids_logprobs_idx.append([token_ids for _ in range(pruned_len)])
pt += pruned_len
return input_token_ids_logprobs_val, input_token_ids_logprobs_idx
@staticmethod
def get_token_ids_logprobs_chunk(
logprobs: torch.Tensor,
logits_metadata: LogitsMetadata,
token_ids_logprobs: List[int],
pruned_lens: List[int],
input_token_ids_logprobs_val: List,
input_token_ids_logprobs_idx: List,
split_pruned_len: int = 0,
):
"""Get token_ids logprobs for each sequence in the chunk.
Args:
logprobs: Log probabilities tensor of shape [seq_len, vocab_size]
logits_metadata: Metadata containing token IDs and pruned length info
token_ids_logprobs: List of token IDs for each sequence
pruned_lens: List of pruned lengths for each sequence
input_token_ids_logprobs_val: List to store token logprob values
input_token_ids_logprobs_idx: List to store token indices
split_pruned_len: Length of pruned tokens from previous chunk
Returns:
int: Number of remaining tokens to process in next chunk
"""
# No sequences in the chunk
if logprobs.shape[0] == 0:
return 0
pt = 0
next_split_pruned_len = 0
for n, (token_ids, pruned_len) in enumerate(
zip(
token_ids_logprobs,
pruned_lens,
)
):
# Adjust pruned length for first sequence
if n == 0:
pruned_len -= split_pruned_len
else:
split_pruned_len = 0
if pruned_len <= 0:
# if pruned length is less than or equal to 0,
# there is no token ids logprobs to process
input_token_ids_logprobs_val.append([])
input_token_ids_logprobs_idx.append([])
continue
# Get the token ids logprobs
val = []
idx = []
for j in range(pruned_len):
# Handle remaining tokens in next chunk if any
if pt + j >= logprobs.shape[0]:
next_split_pruned_len = split_pruned_len + j
break
if token_ids is not None:
val.append(logprobs[pt + j, token_ids].tolist())
idx.append(token_ids)
# Append or extend based on whether the sequence was split across chunks
if len(val) > 0:
if split_pruned_len > 0:
input_token_ids_logprobs_val[-1].extend(val)
input_token_ids_logprobs_idx[-1].extend(idx)
else:
input_token_ids_logprobs_val.append(val)
input_token_ids_logprobs_idx.append(idx)
pt += pruned_len
return next_split_pruned_len
@staticmethod
def compute_temp_top_p_normalized_logprobs(
last_logits: torch.Tensor,
logits_metadata: LogitsMetadata,
top_p: Optional[torch.Tensor] = None,
temperature: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""
compute logprobs for the output token from the given logits.
Returns:
torch.Tensor: logprobs from logits
"""
if top_p is None:
top_p = logits_metadata.top_p
if temperature is None:
temperature = logits_metadata.temperature
# Scale logits if temperature scaling is enabled
if logits_metadata.temp_scaled_logprobs:
last_logits = last_logits / temperature
# Normalize logprobs if top_p normalization is enabled
# NOTE: only normalize logprobs when top_p is set and not equal to 1.0
if logits_metadata.top_p_normalized_logprobs and (top_p != 1.0).any():
from sglang.srt.layers.sampler import top_p_normalize_probs_torch
probs = torch.softmax(last_logits, dim=-1)
del last_logits
probs = top_p_normalize_probs_torch(probs, top_p)
return torch.log(probs)
else:
return torch.nn.functional.log_softmax(last_logits, dim=-1)
@triton.jit
def fused_softcap_kernel(