[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(

View File

@@ -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,

View File

@@ -0,0 +1,2 @@
# Temp workaround, make layer utils more fine-grained later
from sglang.srt.layers.utils.common import *

View 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

View File

@@ -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

View File

@@ -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

View File

@@ -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