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