Move sampler into CUDA graph (#1201)
Co-authored-by: Yineng Zhang <me@zhyncs.com>
This commit is contained in:
@@ -25,16 +25,18 @@ from vllm.distributed.parallel_state import graph_capture
|
||||
from vllm.model_executor.custom_op import CustomOp
|
||||
|
||||
from sglang.srt.layers.logits_processor import (
|
||||
LogitProcessorOutput,
|
||||
LogitsMetadata,
|
||||
LogitsProcessor,
|
||||
LogitsProcessorOutput,
|
||||
)
|
||||
from sglang.srt.layers.sampler import SampleOutput
|
||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||
from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardMode,
|
||||
InputMetadata,
|
||||
update_flashinfer_indices,
|
||||
)
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
from sglang.srt.utils import monkey_patch_vllm_all_gather
|
||||
|
||||
|
||||
@@ -143,6 +145,10 @@ class CudaGraphRunner:
|
||||
self.flashinfer_kv_indices.clone(),
|
||||
]
|
||||
|
||||
# Sampling inputs
|
||||
vocab_size = model_runner.model_config.vocab_size
|
||||
self.sampling_info = SamplingBatchInfo.dummy_one(self.max_bs, vocab_size)
|
||||
|
||||
self.compile_bs = [1, 2, 4, 8, 16, 24, 32] if use_torch_compile else []
|
||||
|
||||
if use_torch_compile:
|
||||
@@ -234,6 +240,7 @@ class CudaGraphRunner:
|
||||
def run_once():
|
||||
input_metadata = InputMetadata(
|
||||
forward_mode=ForwardMode.DECODE,
|
||||
sampling_info=self.sampling_info[:bs],
|
||||
batch_size=bs,
|
||||
req_pool_indices=req_pool_indices,
|
||||
seq_lens=seq_lens,
|
||||
@@ -298,27 +305,35 @@ class CudaGraphRunner:
|
||||
self.flashinfer_handlers[bs],
|
||||
)
|
||||
|
||||
# Sampling inputs
|
||||
self.sampling_info.inplace_assign(raw_bs, batch.sampling_info)
|
||||
|
||||
# Replay
|
||||
torch.cuda.synchronize()
|
||||
self.graphs[bs].replay()
|
||||
torch.cuda.synchronize()
|
||||
output = self.output_buffers[bs]
|
||||
sample_output, logits_output = self.output_buffers[bs]
|
||||
|
||||
# Unpad
|
||||
if bs != raw_bs:
|
||||
output = LogitProcessorOutput(
|
||||
next_token_logits=output.next_token_logits[:raw_bs],
|
||||
logits_output = LogitsProcessorOutput(
|
||||
next_token_logits=logits_output.next_token_logits[:raw_bs],
|
||||
next_token_logprobs=None,
|
||||
normalized_prompt_logprobs=None,
|
||||
input_token_logprobs=None,
|
||||
input_top_logprobs=None,
|
||||
output_top_logprobs=None,
|
||||
)
|
||||
sample_output = SampleOutput(
|
||||
sample_output.success[:raw_bs],
|
||||
sample_output.probs[:raw_bs],
|
||||
sample_output.batch_next_token_ids[:raw_bs],
|
||||
)
|
||||
|
||||
# Extract logprobs
|
||||
if batch.return_logprob:
|
||||
output.next_token_logprobs = torch.nn.functional.log_softmax(
|
||||
output.next_token_logits, dim=-1
|
||||
logits_output.next_token_logprobs = torch.nn.functional.log_softmax(
|
||||
logits_output.next_token_logits, dim=-1
|
||||
)
|
||||
return_top_logprob = any(x > 0 for x in batch.top_logprobs_nums)
|
||||
if return_top_logprob:
|
||||
@@ -326,8 +341,8 @@ class CudaGraphRunner:
|
||||
forward_mode=ForwardMode.DECODE,
|
||||
top_logprobs_nums=batch.top_logprobs_nums,
|
||||
)
|
||||
output.output_top_logprobs = LogitsProcessor.get_top_logprobs(
|
||||
output.next_token_logprobs, logits_metadata
|
||||
logits_output.output_top_logprobs = LogitsProcessor.get_top_logprobs(
|
||||
logits_output.next_token_logprobs, logits_metadata
|
||||
)[1]
|
||||
|
||||
return output
|
||||
return sample_output, logits_output
|
||||
|
||||
Reference in New Issue
Block a user