Clean up wrapper in flashinfer backend (#2638)

This commit is contained in:
Lianmin Zheng
2024-12-29 00:45:57 -08:00
committed by GitHub
parent fd34f2da35
commit 3815b23ccb
12 changed files with 197 additions and 94 deletions

View File

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