From 36361adcbf5257899438833ede09acf6c34b0d98 Mon Sep 17 00:00:00 2001 From: Tiwei Bie Date: Mon, 8 Dec 2025 14:12:35 +0800 Subject: [PATCH] [DLLM] Add initial cuda graph support (#14203) --- .../layers/attention/flashinfer_backend.py | 46 ++++++++++++++++++- python/sglang/srt/layers/logits_processor.py | 8 ++++ python/sglang/srt/managers/schedule_batch.py | 4 +- .../srt/model_executor/cuda_graph_runner.py | 18 +++++++- .../srt/model_executor/forward_batch_info.py | 11 ++++- python/sglang/srt/server_args.py | 14 ++++-- 6 files changed, 93 insertions(+), 8 deletions(-) diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index ecebcf76d..29a1ec83a 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -16,6 +16,7 @@ from typing import TYPE_CHECKING, Callable, List, Optional, Union import torch +from sglang.srt.dllm.config import DllmConfig from sglang.srt.environ import envs from sglang.srt.layers.attention.base_attn_backend import AttentionBackend 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 ) - 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 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.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: 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, 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: raise ValueError("Invalid forward mode") diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index 522865765..67adcd572 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -392,6 +392,14 @@ class LogitsProcessor(nn.Module): 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 if ( logits_metadata.forward_mode.is_decode_or_idle() diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index bf1f13d58..2ceabd431 100644 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -1318,7 +1318,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): ), f"Expected {len(self.out_cache_loc)}, got {self.extend_num_tokens}" 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 reqs = self.reqs diff --git a/python/sglang/srt/model_executor/cuda_graph_runner.py b/python/sglang/srt/model_executor/cuda_graph_runner.py index 7cefcfa16..1fd483da3 100644 --- a/python/sglang/srt/model_executor/cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/cuda_graph_runner.py @@ -40,6 +40,7 @@ from sglang.srt.distributed.parallel_state import ( graph_capture, 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.dp_attention import ( DpPaddingMode, @@ -263,6 +264,9 @@ class CudaGraphRunner: 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 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}") @@ -283,6 +287,9 @@ class CudaGraphRunner: self.num_tokens_per_bs = ( 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 model_runner.server_args.enable_return_hidden_states: @@ -299,6 +306,8 @@ class CudaGraphRunner: self.maybe_init_pdmux() self.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 @@ -825,7 +834,14 @@ class CudaGraphRunner: output = self.output_buffers[graph_key] if isinstance(output, 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=( output.hidden_states[: self.raw_num_token] if output.hidden_states is not None diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 3a85e6a7e..d4ed92d0e 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -91,6 +91,9 @@ class ForwardMode(IntEnum): # Split Prefill for PD multiplexing SPLIT_PREFILL = auto() + # Used in diffusion LLM inference + DLLM_EXTEND = auto() + def is_prefill(self): return self.is_extend() @@ -102,6 +105,7 @@ class ForwardMode(IntEnum): or (include_draft_extend_v2 and self == ForwardMode.DRAFT_EXTEND_V2) or self == ForwardMode.TARGET_VERIFY or self == ForwardMode.SPLIT_PREFILL + or self == ForwardMode.DLLM_EXTEND ) def is_context_parallel_extend(self, include_draft_extend_v2: bool = False): @@ -153,6 +157,7 @@ class ForwardMode(IntEnum): self == ForwardMode.DECODE or self == ForwardMode.TARGET_VERIFY or self == ForwardMode.IDLE + or self == ForwardMode.DLLM_EXTEND ) def is_cpu_graph(self): @@ -171,6 +176,9 @@ class ForwardMode(IntEnum): def is_prebuilt(self): return self == ForwardMode.PREBUILT + def is_dllm_extend(self): + return self == ForwardMode.DLLM_EXTEND + @total_ordering class CaptureHiddenMode(IntEnum): @@ -442,8 +450,9 @@ class ForwardBatch: block_size = batch.dllm_config.block_size 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 i in range(block_offset, block_offset + block_size) ], dtype=torch.int32, ).to(device, non_blocking=True) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 8e7753dab..061511043 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -2050,10 +2050,16 @@ class ServerArgs: if self.dllm_algorithm is None: return if not self.disable_cuda_graph: - logger.warning( - "Cuda graph is disabled because of using diffusion LLM inference" - ) - self.disable_cuda_graph = True + if self.cuda_graph_bs != [1]: + logger.warning( + "Cuda graph bs is set to [1] because of using diffusion LLM inference" + ) + 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: logger.warning( "Overlap schedule is disabled because of using diffusion LLM inference"