Clean up wrapper in flashinfer backend (#2638)
This commit is contained in:
@@ -24,7 +24,11 @@ from vllm.distributed import (
|
||||
)
|
||||
|
||||
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.srt.model_executor.forward_batch_info import (
|
||||
CaptureHiddenMode,
|
||||
ForwardBatch,
|
||||
ForwardMode,
|
||||
)
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
@@ -46,6 +50,10 @@ class LogitsProcessorOutput:
|
||||
output_top_logprobs_val: List = None
|
||||
output_top_logprobs_idx: List = None
|
||||
|
||||
# Used by speculative decoding (EAGLE)
|
||||
# The output of transformer layers
|
||||
hidden_states: Optional[torch.Tensor] = None
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class LogitsMetadata:
|
||||
@@ -61,6 +69,8 @@ class LogitsMetadata:
|
||||
extend_logprob_start_lens_cpu: Optional[List[int]] = None
|
||||
extend_logprob_pruned_lens_cpu: Optional[List[int]] = None
|
||||
|
||||
capture_hidden_mode: CaptureHiddenMode = CaptureHiddenMode.NULL
|
||||
|
||||
@classmethod
|
||||
def from_forward_batch(cls, forward_batch: ForwardBatch):
|
||||
extend_logprob_pruned_lens_cpu = None
|
||||
@@ -78,6 +88,11 @@ class LogitsMetadata:
|
||||
else:
|
||||
return_top_logprob = False
|
||||
|
||||
if forward_batch.spec_info:
|
||||
capture_hidden_mode = forward_batch.spec_info.capture_hidden_mode
|
||||
else:
|
||||
capture_hidden_mode = CaptureHiddenMode.NULL
|
||||
|
||||
return cls(
|
||||
forward_mode=forward_batch.forward_mode,
|
||||
top_logprobs_nums=forward_batch.top_logprobs_nums,
|
||||
@@ -87,6 +102,7 @@ class LogitsMetadata:
|
||||
extend_seq_lens_cpu=forward_batch.extend_seq_lens_cpu,
|
||||
extend_logprob_start_lens_cpu=forward_batch.extend_logprob_start_lens_cpu,
|
||||
extend_logprob_pruned_lens_cpu=extend_logprob_pruned_lens_cpu,
|
||||
capture_hidden_mode=capture_hidden_mode,
|
||||
)
|
||||
|
||||
|
||||
@@ -116,7 +132,10 @@ class LogitsProcessor(nn.Module):
|
||||
assert isinstance(logits_metadata, LogitsMetadata)
|
||||
|
||||
# Get the last hidden states and last logits for the next token prediction
|
||||
if logits_metadata.forward_mode.is_decode():
|
||||
if (
|
||||
logits_metadata.forward_mode.is_decode()
|
||||
or logits_metadata.forward_mode.is_target_verify()
|
||||
):
|
||||
last_index = None
|
||||
last_hidden = hidden_states
|
||||
else:
|
||||
@@ -137,6 +156,15 @@ class LogitsProcessor(nn.Module):
|
||||
if not logits_metadata.return_logprob:
|
||||
return LogitsProcessorOutput(
|
||||
next_token_logits=last_logits,
|
||||
hidden_states=(
|
||||
hidden_states
|
||||
if logits_metadata.capture_hidden_mode.is_full()
|
||||
else (
|
||||
last_hidden
|
||||
if logits_metadata.capture_hidden_mode.is_last()
|
||||
else None
|
||||
)
|
||||
),
|
||||
)
|
||||
else:
|
||||
last_logprobs = self.compute_temp_top_p_normalized_logprobs(
|
||||
|
||||
Reference in New Issue
Block a user