[DLLM] Add initial cuda graph support (#14203)
This commit is contained in:
@@ -16,6 +16,7 @@ from typing import TYPE_CHECKING, Callable, List, Optional, Union
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.dllm.config import DllmConfig
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||||
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton
|
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton
|
||||||
@@ -126,7 +127,9 @@ class FlashInferAttnBackend(AttentionBackend):
|
|||||||
model_runner.server_args.multi_item_scoring_delimiter
|
model_runner.server_args.multi_item_scoring_delimiter
|
||||||
)
|
)
|
||||||
|
|
||||||
self.is_dllm_model = model_runner.server_args.dllm_algorithm is not None
|
# FIXME: remove dllm workarounds from flashinfer
|
||||||
|
self.dllm_config = DllmConfig.from_server_args(model_runner.server_args)
|
||||||
|
self.is_dllm_model = self.dllm_config is not None
|
||||||
|
|
||||||
# Parse constants
|
# Parse constants
|
||||||
self.decode_use_tensor_cores = should_use_tensor_core(
|
self.decode_use_tensor_cores = should_use_tensor_core(
|
||||||
@@ -639,6 +642,35 @@ class FlashInferAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
self.prefill_cuda_graph_metadata[bs] = prefill_wrappers
|
self.prefill_cuda_graph_metadata[bs] = prefill_wrappers
|
||||||
self.forward_metadata = PrefillMetadata(prefill_wrappers, False, False)
|
self.forward_metadata = PrefillMetadata(prefill_wrappers, False, False)
|
||||||
|
elif forward_mode.is_dllm_extend():
|
||||||
|
prefill_wrappers = []
|
||||||
|
for i in range(self.num_wrappers):
|
||||||
|
prefill_wrappers.append(
|
||||||
|
BatchPrefillWithPagedKVCacheWrapper(
|
||||||
|
self.workspace_buffer,
|
||||||
|
"NHD",
|
||||||
|
backend="fa2",
|
||||||
|
use_cuda_graph=True,
|
||||||
|
qo_indptr_buf=self.cuda_graph_qo_indptr[i][: bs + 1],
|
||||||
|
paged_kv_indptr_buf=self.kv_indptr[i][: bs + 1],
|
||||||
|
paged_kv_indices_buf=self.cuda_graph_kv_indices[i],
|
||||||
|
paged_kv_last_page_len_buf=self.kv_last_page_len[:bs],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
seq_lens_sum = seq_lens.sum().item()
|
||||||
|
self.indices_updater_prefill.update(
|
||||||
|
req_pool_indices,
|
||||||
|
seq_lens,
|
||||||
|
seq_lens.cpu(), # may add a little overhead in capture stage
|
||||||
|
seq_lens_sum,
|
||||||
|
prefix_lens=seq_lens - self.dllm_config.block_size,
|
||||||
|
prefill_wrappers=prefill_wrappers,
|
||||||
|
use_ragged=True,
|
||||||
|
encoder_lens=encoder_lens,
|
||||||
|
spec_info=None,
|
||||||
|
)
|
||||||
|
self.prefill_cuda_graph_metadata[bs] = prefill_wrappers
|
||||||
|
self.forward_metadata = PrefillMetadata(prefill_wrappers, True, False)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Invalid mode: {forward_mode=}")
|
raise ValueError(f"Invalid mode: {forward_mode=}")
|
||||||
|
|
||||||
@@ -689,6 +721,18 @@ class FlashInferAttnBackend(AttentionBackend):
|
|||||||
encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None,
|
encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None,
|
||||||
spec_info=spec_info,
|
spec_info=spec_info,
|
||||||
)
|
)
|
||||||
|
elif forward_mode.is_dllm_extend():
|
||||||
|
self.indices_updater_prefill.update(
|
||||||
|
req_pool_indices[:bs],
|
||||||
|
seq_lens[:bs],
|
||||||
|
seq_lens_cpu[:bs] if seq_lens_cpu is not None else None,
|
||||||
|
seq_lens_sum,
|
||||||
|
prefix_lens=seq_lens - self.dllm_config.block_size,
|
||||||
|
prefill_wrappers=self.prefill_cuda_graph_metadata[bs],
|
||||||
|
use_ragged=True,
|
||||||
|
encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None,
|
||||||
|
spec_info=None,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
raise ValueError("Invalid forward mode")
|
raise ValueError("Invalid forward mode")
|
||||||
|
|
||||||
|
|||||||
@@ -392,6 +392,14 @@ class LogitsProcessor(nn.Module):
|
|||||||
input_ids, hidden_states, lm_head, logits_metadata, multi_item_delimiter
|
input_ids, hidden_states, lm_head, logits_metadata, multi_item_delimiter
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if logits_metadata.forward_mode.is_dllm_extend():
|
||||||
|
assert self.return_full_logits
|
||||||
|
full_logits = self._get_logits(hidden_states, lm_head, logits_metadata)
|
||||||
|
return LogitsProcessorOutput(
|
||||||
|
full_logits=full_logits,
|
||||||
|
next_token_logits=None,
|
||||||
|
)
|
||||||
|
|
||||||
# Get the last hidden states and last logits for the next token prediction
|
# Get the last hidden states and last logits for the next token prediction
|
||||||
if (
|
if (
|
||||||
logits_metadata.forward_mode.is_decode_or_idle()
|
logits_metadata.forward_mode.is_decode_or_idle()
|
||||||
|
|||||||
@@ -1318,7 +1318,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
), f"Expected {len(self.out_cache_loc)}, got {self.extend_num_tokens}"
|
), f"Expected {len(self.out_cache_loc)}, got {self.extend_num_tokens}"
|
||||||
|
|
||||||
def prepare_for_extend(self):
|
def prepare_for_extend(self):
|
||||||
self.forward_mode = ForwardMode.EXTEND
|
self.forward_mode = (
|
||||||
|
ForwardMode.DLLM_EXTEND if self.is_dllm() else ForwardMode.EXTEND
|
||||||
|
)
|
||||||
|
|
||||||
# Init tensors
|
# Init tensors
|
||||||
reqs = self.reqs
|
reqs = self.reqs
|
||||||
|
|||||||
@@ -40,6 +40,7 @@ from sglang.srt.distributed.parallel_state import (
|
|||||||
graph_capture,
|
graph_capture,
|
||||||
set_pdmux_status,
|
set_pdmux_status,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.dllm.config import DllmConfig
|
||||||
from sglang.srt.layers.attention.nsa.utils import is_nsa_enable_prefill_cp
|
from sglang.srt.layers.attention.nsa.utils import is_nsa_enable_prefill_cp
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
DpPaddingMode,
|
DpPaddingMode,
|
||||||
@@ -263,6 +264,9 @@ class CudaGraphRunner:
|
|||||||
|
|
||||||
self.deepep_adapter = DeepEPCudaGraphRunnerAdapter()
|
self.deepep_adapter = DeepEPCudaGraphRunnerAdapter()
|
||||||
|
|
||||||
|
self.dllm_config = DllmConfig.from_server_args(model_runner.server_args)
|
||||||
|
self.is_dllm = self.dllm_config is not None
|
||||||
|
|
||||||
# Batch sizes to capture
|
# Batch sizes to capture
|
||||||
self.capture_bs, self.compile_bs = get_batch_sizes_to_capture(model_runner)
|
self.capture_bs, self.compile_bs = get_batch_sizes_to_capture(model_runner)
|
||||||
log_info_on_rank0(logger, f"Capture cuda graph bs {self.capture_bs}")
|
log_info_on_rank0(logger, f"Capture cuda graph bs {self.capture_bs}")
|
||||||
@@ -283,6 +287,9 @@ class CudaGraphRunner:
|
|||||||
self.num_tokens_per_bs = (
|
self.num_tokens_per_bs = (
|
||||||
self.model_runner.server_args.speculative_num_draft_tokens
|
self.model_runner.server_args.speculative_num_draft_tokens
|
||||||
)
|
)
|
||||||
|
elif self.is_dllm:
|
||||||
|
self.capture_forward_mode = ForwardMode.DLLM_EXTEND
|
||||||
|
self.num_tokens_per_bs = self.dllm_config.block_size
|
||||||
|
|
||||||
# If returning hidden states is enabled, set initial capture hidden mode to full to avoid double-capture on startup
|
# If returning hidden states is enabled, set initial capture hidden mode to full to avoid double-capture on startup
|
||||||
if model_runner.server_args.enable_return_hidden_states:
|
if model_runner.server_args.enable_return_hidden_states:
|
||||||
@@ -299,6 +306,8 @@ class CudaGraphRunner:
|
|||||||
self.maybe_init_pdmux()
|
self.maybe_init_pdmux()
|
||||||
self.seq_len_fill_value = (
|
self.seq_len_fill_value = (
|
||||||
self.model_runner.attn_backend.get_cuda_graph_seq_len_fill_value()
|
self.model_runner.attn_backend.get_cuda_graph_seq_len_fill_value()
|
||||||
|
if self.dllm_config is None
|
||||||
|
else self.dllm_config.block_size
|
||||||
)
|
)
|
||||||
|
|
||||||
self.encoder_len_fill_value = 0
|
self.encoder_len_fill_value = 0
|
||||||
@@ -825,7 +834,14 @@ class CudaGraphRunner:
|
|||||||
output = self.output_buffers[graph_key]
|
output = self.output_buffers[graph_key]
|
||||||
if isinstance(output, LogitsProcessorOutput):
|
if isinstance(output, LogitsProcessorOutput):
|
||||||
return LogitsProcessorOutput(
|
return LogitsProcessorOutput(
|
||||||
next_token_logits=output.next_token_logits[: self.raw_num_token],
|
next_token_logits=(
|
||||||
|
output.next_token_logits[: self.raw_num_token]
|
||||||
|
if not self.is_dllm
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
full_logits=(
|
||||||
|
output.full_logits[: self.raw_num_token] if self.is_dllm else None
|
||||||
|
),
|
||||||
hidden_states=(
|
hidden_states=(
|
||||||
output.hidden_states[: self.raw_num_token]
|
output.hidden_states[: self.raw_num_token]
|
||||||
if output.hidden_states is not None
|
if output.hidden_states is not None
|
||||||
|
|||||||
@@ -91,6 +91,9 @@ class ForwardMode(IntEnum):
|
|||||||
# Split Prefill for PD multiplexing
|
# Split Prefill for PD multiplexing
|
||||||
SPLIT_PREFILL = auto()
|
SPLIT_PREFILL = auto()
|
||||||
|
|
||||||
|
# Used in diffusion LLM inference
|
||||||
|
DLLM_EXTEND = auto()
|
||||||
|
|
||||||
def is_prefill(self):
|
def is_prefill(self):
|
||||||
return self.is_extend()
|
return self.is_extend()
|
||||||
|
|
||||||
@@ -102,6 +105,7 @@ class ForwardMode(IntEnum):
|
|||||||
or (include_draft_extend_v2 and self == ForwardMode.DRAFT_EXTEND_V2)
|
or (include_draft_extend_v2 and self == ForwardMode.DRAFT_EXTEND_V2)
|
||||||
or self == ForwardMode.TARGET_VERIFY
|
or self == ForwardMode.TARGET_VERIFY
|
||||||
or self == ForwardMode.SPLIT_PREFILL
|
or self == ForwardMode.SPLIT_PREFILL
|
||||||
|
or self == ForwardMode.DLLM_EXTEND
|
||||||
)
|
)
|
||||||
|
|
||||||
def is_context_parallel_extend(self, include_draft_extend_v2: bool = False):
|
def is_context_parallel_extend(self, include_draft_extend_v2: bool = False):
|
||||||
@@ -153,6 +157,7 @@ class ForwardMode(IntEnum):
|
|||||||
self == ForwardMode.DECODE
|
self == ForwardMode.DECODE
|
||||||
or self == ForwardMode.TARGET_VERIFY
|
or self == ForwardMode.TARGET_VERIFY
|
||||||
or self == ForwardMode.IDLE
|
or self == ForwardMode.IDLE
|
||||||
|
or self == ForwardMode.DLLM_EXTEND
|
||||||
)
|
)
|
||||||
|
|
||||||
def is_cpu_graph(self):
|
def is_cpu_graph(self):
|
||||||
@@ -171,6 +176,9 @@ class ForwardMode(IntEnum):
|
|||||||
def is_prebuilt(self):
|
def is_prebuilt(self):
|
||||||
return self == ForwardMode.PREBUILT
|
return self == ForwardMode.PREBUILT
|
||||||
|
|
||||||
|
def is_dllm_extend(self):
|
||||||
|
return self == ForwardMode.DLLM_EXTEND
|
||||||
|
|
||||||
|
|
||||||
@total_ordering
|
@total_ordering
|
||||||
class CaptureHiddenMode(IntEnum):
|
class CaptureHiddenMode(IntEnum):
|
||||||
@@ -442,8 +450,9 @@ class ForwardBatch:
|
|||||||
block_size = batch.dllm_config.block_size
|
block_size = batch.dllm_config.block_size
|
||||||
ret.positions = torch.tensor(
|
ret.positions = torch.tensor(
|
||||||
[
|
[
|
||||||
[i for i in range(block_offset, block_offset + block_size)]
|
i
|
||||||
for block_offset in batch.dllm_block_offsets
|
for block_offset in batch.dllm_block_offsets
|
||||||
|
for i in range(block_offset, block_offset + block_size)
|
||||||
],
|
],
|
||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
).to(device, non_blocking=True)
|
).to(device, non_blocking=True)
|
||||||
|
|||||||
@@ -2050,10 +2050,16 @@ class ServerArgs:
|
|||||||
if self.dllm_algorithm is None:
|
if self.dllm_algorithm is None:
|
||||||
return
|
return
|
||||||
if not self.disable_cuda_graph:
|
if not self.disable_cuda_graph:
|
||||||
logger.warning(
|
if self.cuda_graph_bs != [1]:
|
||||||
"Cuda graph is disabled because of using diffusion LLM inference"
|
logger.warning(
|
||||||
)
|
"Cuda graph bs is set to [1] because of using diffusion LLM inference"
|
||||||
self.disable_cuda_graph = True
|
)
|
||||||
|
self.cuda_graph_bs = [1]
|
||||||
|
if self.attention_backend != "flashinfer":
|
||||||
|
logger.warning(
|
||||||
|
"Attention backend is set to flashinfer because of enabling cuda graph in diffusion LLM inference"
|
||||||
|
)
|
||||||
|
self.attention_backend = "flashinfer"
|
||||||
if not self.disable_overlap_schedule:
|
if not self.disable_overlap_schedule:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Overlap schedule is disabled because of using diffusion LLM inference"
|
"Overlap schedule is disabled because of using diffusion LLM inference"
|
||||||
|
|||||||
Reference in New Issue
Block a user