Minor follow-up fixes for the logprob refactor (#2670)

This commit is contained in:
Lianmin Zheng
2024-12-30 05:42:08 -08:00
committed by GitHub
parent c5210dfa38
commit 21ec66e59e
5 changed files with 11 additions and 12 deletions
+3 -3
View File
@@ -35,21 +35,21 @@ from sglang.srt.model_executor.forward_batch_info import (
@dataclasses.dataclass
class LogitsProcessorOutput:
## First part. This part will be returned by python/sglang/srt/layers/logits_processor.py::LogitsProcessor.
## Part 1: This part will be assigned in python/sglang/srt/layers/logits_processor.py::LogitsProcessor
# The logits of the next tokens. shape: [#seq, vocab_size]
next_token_logits: torch.Tensor
# Used by speculative decoding (EAGLE)
# The last hidden layers
hidden_states: Optional[torch.Tensor] = None
## Second part. This part will be returned by python/sglang/srt/layers/sampler.py::Sampler.
## Part 2: This part will be assigned in python/sglang/srt/layers/sampler.py::Sampler
# The logprobs of the next tokens. shape: [#seq]
next_token_logprobs: Optional[torch.Tensor] = None
# The logprobs and ids of the top-k tokens in output positions. shape: [#seq, k]
next_token_top_logprobs_val: Optional[List] = None
next_token_top_logprobs_idx: Optional[List] = None
## Third part. This part will be returned by python/sglang/srt/layers/logits_processor.py::LogitsProcessor. Prefill-only.
## Part 3: Prefill-only. This part will be assigned in python/sglang/srt/layers/logits_processor.py::LogitsProcessor
# The normlaized logprobs of prompts. shape: [#seq]
normalized_prompt_logprobs: torch.Tensor = None
# The logprobs of input tokens. shape: [#token]
+3 -1
View File
@@ -56,7 +56,9 @@ class Sampler(nn.Module):
if global_server_args_dict["sampling_backend"] == "flashinfer":
if return_logprob:
# NOTE: the top_p_renorm_prob from flashinfer has numerical problems
# NOTE: the top_p_renorm_prob from flashinfer has numerical problems,
# https://github.com/flashinfer-ai/flashinfer/issues/708
# so we use the torch implementation.
logprobs = torch.log(
top_p_normalize_probs_torch(probs, sampling_info.top_ps)
)