diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index 9426c108f..293e69c42 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -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( diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py index 4e22d1f83..55bef5652 100644 --- a/python/sglang/srt/layers/sampler.py +++ b/python/sglang/srt/layers/sampler.py @@ -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, diff --git a/python/sglang/srt/layers/utils/__init__.py b/python/sglang/srt/layers/utils/__init__.py new file mode 100644 index 000000000..7881b2b49 --- /dev/null +++ b/python/sglang/srt/layers/utils/__init__.py @@ -0,0 +1,2 @@ +# Temp workaround, make layer utils more fine-grained later +from sglang.srt.layers.utils.common import * diff --git a/python/sglang/srt/layers/utils.py b/python/sglang/srt/layers/utils/common.py similarity index 100% rename from python/sglang/srt/layers/utils.py rename to python/sglang/srt/layers/utils/common.py diff --git a/python/sglang/srt/layers/utils/logprob.py b/python/sglang/srt/layers/utils/logprob.py new file mode 100644 index 000000000..ebc3031da --- /dev/null +++ b/python/sglang/srt/layers/utils/logprob.py @@ -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 diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index f681ca158..ed9e752c5 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -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 diff --git a/python/sglang/srt/speculative/mtp_worker.py b/python/sglang/srt/speculative/mtp_worker.py index 233fd1992..24cd20a98 100644 --- a/python/sglang/srt/speculative/mtp_worker.py +++ b/python/sglang/srt/speculative/mtp_worker.py @@ -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 diff --git a/python/sglang/srt/speculative/ngram_worker.py b/python/sglang/srt/speculative/ngram_worker.py index 5c61ab31a..4296c6924 100644 --- a/python/sglang/srt/speculative/ngram_worker.py +++ b/python/sglang/srt/speculative/ngram_worker.py @@ -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