[1 / N] Clean up logprob utils (#15509)
This commit is contained in:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user