[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(
|
||||
|
||||
@@ -11,6 +11,7 @@ from sglang.srt.layers.dp_attention import (
|
||||
is_dp_attention_enabled,
|
||||
)
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
from sglang.srt.layers.utils.logprob import get_token_ids_logprobs, get_top_logprobs
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
from sglang.srt.sampling.sampling_params import TOP_K_ALL
|
||||
from sglang.srt.server_args import get_global_server_args
|
||||
@@ -452,27 +453,6 @@ def top_p_normalize_probs_torch(
|
||||
return torch.zeros_like(probs_sort).scatter_(-1, probs_idx, probs_sort)
|
||||
|
||||
|
||||
def get_top_logprobs(
|
||||
logprobs: torch.Tensor,
|
||||
top_logprobs_nums: List[int],
|
||||
):
|
||||
max_k = max(top_logprobs_nums)
|
||||
ret = logprobs.topk(max_k, dim=1)
|
||||
values = ret.values.tolist()
|
||||
indices = ret.indices.tolist()
|
||||
|
||||
output_top_logprobs_val = []
|
||||
output_top_logprobs_idx = []
|
||||
for i, k in enumerate(top_logprobs_nums):
|
||||
output_top_logprobs_val.append(values[i][:k])
|
||||
output_top_logprobs_idx.append(indices[i][:k])
|
||||
|
||||
return (
|
||||
output_top_logprobs_val,
|
||||
output_top_logprobs_idx,
|
||||
)
|
||||
|
||||
|
||||
def get_token_ids_logprobs_batch_optimized(
|
||||
logprobs: torch.Tensor,
|
||||
token_ids_logprobs: List[List[int]],
|
||||
@@ -561,23 +541,6 @@ def get_token_ids_logprobs_batch_optimized(
|
||||
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):
|
||||
if token_ids is not None:
|
||||
output_token_ids_logprobs_val.append(logprobs[i, token_ids].tolist())
|
||||
output_token_ids_logprobs_idx.append(token_ids)
|
||||
else:
|
||||
output_token_ids_logprobs_val.append([])
|
||||
output_token_ids_logprobs_idx.append([])
|
||||
|
||||
return (
|
||||
output_token_ids_logprobs_val,
|
||||
output_token_ids_logprobs_idx,
|
||||
)
|
||||
|
||||
|
||||
def apply_custom_logit_processor(
|
||||
logits: torch.Tensor,
|
||||
sampling_batch_info: SamplingBatchInfo,
|
||||
|
||||
2
python/sglang/srt/layers/utils/__init__.py
Normal file
2
python/sglang/srt/layers/utils/__init__.py
Normal file
@@ -0,0 +1,2 @@
|
||||
# Temp workaround, make layer utils more fine-grained later
|
||||
from sglang.srt.layers.utils.common import *
|
||||
306
python/sglang/srt/layers/utils/logprob.py
Normal file
306
python/sglang/srt/layers/utils/logprob.py
Normal file
@@ -0,0 +1,306 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
from enum import Enum, auto
|
||||
from typing import TYPE_CHECKING, List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.logits_processor import LogitsMetadata
|
||||
|
||||
|
||||
class LogprobStage(Enum):
|
||||
PREFILL = auto()
|
||||
DECODE = auto()
|
||||
|
||||
|
||||
@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
|
||||
|
||||
|
||||
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)
|
||||
|
||||
|
||||
def get_top_logprobs_raw(
|
||||
logprobs: torch.Tensor,
|
||||
top_logprobs_nums: List[int],
|
||||
stage: LogprobStage,
|
||||
extend_logprob_pruned_lens_cpu: Optional[List[int]] = None,
|
||||
):
|
||||
max_k = max(top_logprobs_nums)
|
||||
values, indices = logprobs.topk(max_k, dim=-1)
|
||||
values = values.tolist()
|
||||
indices = indices.tolist()
|
||||
|
||||
top_logprobs_val = []
|
||||
top_logprobs_idx = []
|
||||
|
||||
if stage == LogprobStage.DECODE:
|
||||
for i, k in enumerate(top_logprobs_nums):
|
||||
top_logprobs_val.append(values[i][:k])
|
||||
top_logprobs_idx.append(indices[i][:k])
|
||||
else:
|
||||
pt = 0
|
||||
for k, pruned_len in zip(top_logprobs_nums, extend_logprob_pruned_lens_cpu):
|
||||
if pruned_len <= 0:
|
||||
top_logprobs_val.append([])
|
||||
top_logprobs_idx.append([])
|
||||
continue
|
||||
|
||||
top_logprobs_val.append([values[pt + j][:k] for j in range(pruned_len)])
|
||||
top_logprobs_idx.append([indices[pt + j][:k] for j in range(pruned_len)])
|
||||
pt += pruned_len
|
||||
|
||||
return top_logprobs_val, top_logprobs_idx
|
||||
|
||||
|
||||
def get_top_logprobs_prefill(
|
||||
all_logprobs: torch.Tensor, logits_metadata: LogitsMetadata
|
||||
):
|
||||
return get_top_logprobs_raw(
|
||||
all_logprobs,
|
||||
logits_metadata.top_logprobs_nums,
|
||||
stage=LogprobStage.PREFILL,
|
||||
extend_logprob_pruned_lens_cpu=logits_metadata.extend_logprob_pruned_lens_cpu,
|
||||
)
|
||||
|
||||
|
||||
def get_top_logprobs(
|
||||
logprobs: torch.Tensor,
|
||||
top_logprobs_nums: List[int],
|
||||
):
|
||||
return get_top_logprobs_raw(logprobs, top_logprobs_nums, stage=LogprobStage.DECODE)
|
||||
|
||||
|
||||
def get_token_ids_logprobs_raw(
|
||||
logprobs: torch.Tensor,
|
||||
token_ids_logprobs: List[Optional[List[int]]],
|
||||
stage: LogprobStage,
|
||||
extend_logprob_pruned_lens_cpu: Optional[List[int]] = None,
|
||||
delay_cpu_copy: bool = False,
|
||||
):
|
||||
vals, idxs = [], []
|
||||
if stage == LogprobStage.DECODE:
|
||||
for i, token_ids in enumerate(token_ids_logprobs):
|
||||
if token_ids is None:
|
||||
vals.append([])
|
||||
idxs.append([])
|
||||
else:
|
||||
vals.append(logprobs[i, token_ids].tolist())
|
||||
idxs.append(token_ids)
|
||||
else: # prefill
|
||||
pt = 0
|
||||
for token_ids, pruned_len in zip(
|
||||
token_ids_logprobs, extend_logprob_pruned_lens_cpu
|
||||
):
|
||||
if pruned_len <= 0:
|
||||
vals.append([])
|
||||
idxs.append([])
|
||||
continue
|
||||
pos_logprobs = logprobs[pt : pt + pruned_len, token_ids]
|
||||
vals.append(pos_logprobs if delay_cpu_copy else pos_logprobs.tolist())
|
||||
idxs.append([token_ids for _ in range(pruned_len)])
|
||||
pt += pruned_len
|
||||
return vals, idxs
|
||||
|
||||
|
||||
def get_token_ids_logprobs_prefill(
|
||||
all_logprobs, logits_metadata: LogitsMetadata, delay_cpu_copy=False
|
||||
):
|
||||
return get_token_ids_logprobs_raw(
|
||||
all_logprobs,
|
||||
logits_metadata.token_ids_logprobs,
|
||||
stage=LogprobStage.PREFILL,
|
||||
extend_logprob_pruned_lens_cpu=logits_metadata.extend_logprob_pruned_lens_cpu,
|
||||
delay_cpu_copy=delay_cpu_copy,
|
||||
)
|
||||
|
||||
|
||||
def get_token_ids_logprobs(logprobs, token_ids_logprobs):
|
||||
return get_token_ids_logprobs_raw(
|
||||
logprobs, token_ids_logprobs, stage=LogprobStage.DECODE
|
||||
)
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
def get_token_ids_logprobs_chunk(
|
||||
logprobs: torch.Tensor,
|
||||
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
|
||||
@@ -14,7 +14,7 @@ from sglang.srt.layers.moe.utils import (
|
||||
speculative_moe_a2a_backend_context,
|
||||
speculative_moe_backend_context,
|
||||
)
|
||||
from sglang.srt.layers.sampler import get_token_ids_logprobs, get_top_logprobs
|
||||
from sglang.srt.layers.utils.logprob import get_token_ids_logprobs, get_top_logprobs
|
||||
from sglang.srt.managers.io_struct import UpdateWeightsFromTensorReqInput
|
||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||
from sglang.srt.managers.scheduler import GenerationBatchResult
|
||||
|
||||
@@ -22,7 +22,7 @@ from sglang.srt.distributed import get_tp_group
|
||||
from sglang.srt.layers.dp_attention import get_attention_tp_group
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
from sglang.srt.layers.moe.utils import speculative_moe_backend_context
|
||||
from sglang.srt.layers.sampler import get_token_ids_logprobs, get_top_logprobs
|
||||
from sglang.srt.layers.utils.logprob import get_token_ids_logprobs, get_top_logprobs
|
||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||
from sglang.srt.managers.scheduler import GenerationBatchResult
|
||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||
|
||||
@@ -7,7 +7,7 @@ from sgl_kernel.speculative import reconstruct_indices_from_tree_mask
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
from sglang.srt.layers.sampler import get_token_ids_logprobs, get_top_logprobs
|
||||
from sglang.srt.layers.utils.logprob import get_token_ids_logprobs, get_top_logprobs
|
||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||
from sglang.srt.managers.scheduler import GenerationBatchResult
|
||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||
|
||||
Reference in New Issue
Block a user