diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py index 77bde4232..690a8535a 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py @@ -469,33 +469,15 @@ class AscendAttnBackend(AttentionBackend): device=self.device, ) - def init_forward_metadata_capture_cuda_graph( + def _init_cuda_graph_metadata( self, bs: int, - num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ): + seq_lens: torch.Tensor, + ) -> "ForwardMetadata": + """Create and store the per-bs ForwardMetadata for CUDA graph capture.""" metadata = ForwardMetadata() - metadata.block_tables = self.graph_metadata["block_tables"][:bs, :] - if self.is_dllm_model: - max_len = int(seq_lens[:bs].max().item()) - max_seq_pages = (max_len + self.page_size - 1) // self.page_size - metadata.block_tables[:bs, :max_seq_pages].copy_( - ( - self.req_to_token[req_pool_indices[:bs], :max_len][ - :, :: self.page_size - ] - // self.page_size - ).to(torch.int32) - ) - metadata.block_tables[:bs, max_seq_pages:].fill_(0) - metadata.block_tables[bs:, :].fill_(0) - if self.is_hybrid_swa: metadata.block_tables_swa = self.graph_metadata["block_tables_swa"][:bs, :] metadata.seq_lens_cpu_list = seq_lens.cpu().int().tolist() @@ -515,7 +497,7 @@ class AscendAttnBackend(AttentionBackend): ) else: metadata.actual_seq_lengths_q = torch.tensor( - [1 + i * 1 for i in range(bs)], + [1 + i for i in range(bs)], dtype=torch.int32, device=seq_lens.device, ) @@ -528,13 +510,11 @@ class AscendAttnBackend(AttentionBackend): metadata.seq_lens_list_cumsum = ( torch.cumsum(extend_seq_lens_cpu_int, dim=0).int().tolist() ) - if ( self.q_head_num_padding is not None and self.q_head_num_padding > self.tp_q_head_num ): - # In the MLA architecture, the FIA kernel requires the head count to be a power of 2. - # Therefore, we pad the head dimension accordingly and initialize an empty tensor for padding. + dtype = self.model_dtype if self.model_dtype is not None else torch.bfloat16 metadata.nope_padding = torch.empty( [ bs, @@ -542,9 +522,7 @@ class AscendAttnBackend(AttentionBackend): self.q_head_num_padding - self.tp_q_head_num, self.kv_lora_rank, ], - dtype=( - self.model_dtype if self.model_dtype is not None else torch.bfloat16 - ), + dtype=dtype, device=seq_lens.device, ) metadata.rope_padding = torch.empty( @@ -554,16 +532,33 @@ class AscendAttnBackend(AttentionBackend): self.q_head_num_padding - self.tp_q_head_num, self.qk_rope_head_dim, ], - dtype=( - self.model_dtype if self.model_dtype is not None else torch.bfloat16 - ), + dtype=dtype, device=seq_lens.device, ) - self.graph_metadata[bs] = metadata - self.forward_metadata = metadata + return metadata - self.graph_mode = True + def init_forward_metadata_capture_cuda_graph( + self, + bs: int, + num_tokens: int, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + encoder_lens: Optional[torch.Tensor], + forward_mode: ForwardMode, + spec_info: Optional[SpecInput], + ): + self._init_cuda_graph_metadata(bs, forward_mode, seq_lens) + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=None, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=seq_lens.cpu(), + ) def init_forward_metadata_replay_cuda_graph( self, diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_gdn_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_gdn_backend.py index e15d3d239..109d474c1 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_gdn_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_gdn_backend.py @@ -93,19 +93,16 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase): forward_mode: ForwardMode, spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], ): - if forward_mode.is_draft_extend(True): - return - super().init_forward_metadata_capture_cuda_graph( - bs, - num_tokens, - req_pool_indices, - seq_lens, - encoder_lens, - forward_mode, - spec_info, + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=None, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=seq_lens.cpu(), ) - self.prepare_gdn_inputs(bs, forward_mode, spec_info) - self.graph_mode = True def init_forward_metadata_replay_cuda_graph( self, diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 05ec7b1cd..9b0fcbf42 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -1492,423 +1492,16 @@ class AiterAttnBackend(AttentionBackend): forward_mode: ForwardMode, spec_info: Optional[SpecInput], ): - - num_kv_splits = None - # num_kv_splits_indptr = None - - work_metadata = None - work_info_set = None - work_indptr = None - - reduce_indptr = None - reduce_final_map = None - reduce_partial_map = None - - swa_page_table = None - - max_kv_len = torch.max(seq_lens).item() - - if forward_mode.is_decode_or_idle(): - qo_indptr = None - kv_last_page_len = None - max_q_len = None - - if spec_info is None or ( - self.use_triton_unified_attention and not self.use_mla - ): - max_num_blocks_per_seq = ( - self.max_context_len + self.page_size - 1 - ) // self.page_size - - if not self.use_triton_unified_attention: - kv_indptr = self.kv_indptr - kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0) - kv_indptr = kv_indptr[: bs + 1] - kv_indices = self.cuda_graph_kv_indices - create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices, - seq_lens, - kv_indptr, - None, - kv_indices, - self.req_to_token.stride(0), - ) - else: - max_q_len = 1 - kv_indices = self.cuda_graph_page_table - - if self.use_sliding_window_kv_pool: - swa_page_table = self.cuda_graph_swa_page_table - - if spec_info is not None: - self._build_unified_page_table_from_spec( - spec_info, - bs, - dest_buf=kv_indices, - swa_dest_buf=swa_page_table, - ) - else: - page_indices = self.req_to_token[ - req_pool_indices[:bs], :max_kv_len - ] - - if self.use_sliding_window_kv_pool: - swa_page_indices = ( - self.token_to_kv_pool.translate_loc_from_full_to_swa( - page_indices - ) - ) - - page_indices = self._transform_table_1_to_real(page_indices) - swa_page_indices = self._transform_table_1_to_real( - swa_page_indices - ) - - new_rows = swa_page_indices.shape[0] - new_cols = swa_page_indices.shape[1] - - kv_indices[:new_rows, :new_cols].copy_(page_indices) - swa_page_table = self.cuda_graph_swa_page_table - swa_page_table[:new_rows, :new_cols].copy_(swa_page_indices) - elif self.page_size > 1: - page_indices = self._transform_table_1_to_real(page_indices) - new_rows = page_indices.shape[0] - new_cols = page_indices.shape[1] - kv_indices[:new_rows, :new_cols].copy_(page_indices) - - qo_indptr = self.qo_indptr_unified_decode[: bs + 1] - - kv_indptr = None - else: - kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices - - if self.use_mla: - qo_indptr = self.qo_indptr_[: bs + 1] - qo_indptr[1 : bs + 1] = torch.cumsum( - self.cuda_graph_kv_last_page_len[:bs], dim=0 - ) - kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs] - max_q_len = 1 - - if _use_mla_ps_kernel: - num_kv_splits = self.max_split_per_batch - - self.make_mla_meta_data( - qo_indptr, - kv_indptr, - kv_last_page_len, - self.work_metadata, - self.work_info_set, - self.work_indptr, - self.reduce_indptr, - self.reduce_final_map, - self.reduce_partial_map, - max_q_len, - fast_mode=fast_mode, - max_split_per_batch=num_kv_splits, - intra_batch_mode=intra_batch_mode, - ) - - work_metadata = self.work_metadata - work_info_set = self.work_info_set - work_indptr = self.work_indptr - - reduce_indptr = self.reduce_indptr - reduce_final_map = self.reduce_final_map - reduce_partial_map = self.reduce_partial_map - - self.forward_metadata = ForwardMetadata( - kv_indptr, - kv_indices, - qo_indptr, - kv_last_page_len, - max_q_len, - max_kv_len, - work_metadata=work_metadata, - work_info_set=work_info_set, - work_indptr=work_indptr, - reduce_indptr=reduce_indptr, - reduce_final_map=reduce_final_map, - reduce_partial_map=reduce_partial_map, - num_kv_splits=num_kv_splits, - swa_page_table=swa_page_table, - ) - - elif forward_mode.is_target_verify(): - qo_indptr = self.qo_indptr[: bs + 1] - qo_indptr[: bs + 1] = torch.arange( - 0, - (1 + bs) * self.num_draft_tokens, - step=self.num_draft_tokens, - dtype=torch.int32, - device=self.device, - ) - if self.use_mla: - kv_lens = seq_lens + self.num_draft_tokens - else: - kv_lens = seq_lens - kv_indptr = self.kv_indptr[: bs + 1] - kv_indptr[1 : bs + 1] = torch.cumsum(kv_lens, dim=0) - kv_indices = self.cuda_graph_kv_indices - create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices, - kv_lens, - kv_indptr, - None, - kv_indices, - self.req_to_token.stride(0), - ) - kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs] - max_q_len = self.num_draft_tokens - - if self.use_mla: - if _use_mla_ps_kernel: - num_kv_splits = self.max_split_per_batch - - self.make_mla_meta_data( - qo_indptr, - kv_indptr, - kv_last_page_len, - self.work_metadata, - self.work_info_set, - self.work_indptr, - self.reduce_indptr, - self.reduce_final_map, - self.reduce_partial_map, - max_q_len, - fast_mode=fast_mode, - max_split_per_batch=num_kv_splits, - intra_batch_mode=intra_batch_mode, - ) - - work_metadata = self.work_metadata - work_info_set = self.work_info_set - work_indptr = self.work_indptr - - reduce_indptr = self.reduce_indptr - reduce_final_map = self.reduce_final_map - reduce_partial_map = self.reduce_partial_map - - self.forward_metadata = ForwardMetadata( - kv_indptr, - kv_indices, - qo_indptr, - kv_last_page_len, - max_q_len, - max_kv_len, - work_metadata=work_metadata, - work_info_set=work_info_set, - work_indptr=work_indptr, - reduce_indptr=reduce_indptr, - reduce_final_map=reduce_final_map, - reduce_partial_map=reduce_partial_map, - num_kv_splits=num_kv_splits, - ) - else: - if self._use_unified_verify: - max_num_blocks_per_seq = ( - self.max_context_len + self.page_size - 1 - ) // self.page_size - page_table = self.cuda_graph_page_table[:bs] - - swa_page_table = None - - if self.use_sliding_window_kv_pool: - swa_page_table = self.cuda_graph_swa_page_table.view( - -1, max_num_blocks_per_seq - )[:bs] - - _page_table, _qo_indptr, _max_q_len, _swa_page_table = ( - self._build_verify_unified_metadata( - bs, - seq_lens, - req_pool_indices, - self.num_draft_tokens, - page_table_dest=page_table, - swa_page_table_dest=swa_page_table, - ) - ) - max_kv_len = max_num_blocks_per_seq * self.page_size - self.forward_metadata = ForwardMetadata( - None, - _page_table, - _qo_indptr, - kv_last_page_len, - _max_q_len, - max_kv_len, - max_extend_len=_max_q_len, - swa_page_table=_swa_page_table, - ) - else: - custom_mask = self.cuda_graph_custom_mask - custom_mask[: spec_info.custom_mask.shape[0]] = ( - spec_info.custom_mask - ) - seq_mask_len = max_q_len * (seq_lens + max_q_len) - mask_indptr = self.mask_indptr - mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len[:bs], dim=0) - mask_indptr = mask_indptr[: bs + 1] - - self.forward_metadata = ForwardMetadata( - kv_indptr, - kv_indices, - qo_indptr, - kv_last_page_len, - max_q_len, - max_kv_len, - custom_mask=custom_mask, - mask_indptr=mask_indptr, - max_extend_len=max_q_len, - ) - elif forward_mode.is_draft_extend_v2(): - # EAGLE V2: Uses fixed num_draft_tokens per batch - self._ensure_spec_v2_topk_supported() - num_tokens_per_bs = self._resolve_v2_num_draft_tokens() - qo_indptr = self._set_uniform_qo_indptr(bs, num_tokens_per_bs, self.device) - kv_indptr = self.kv_indptr[: bs + 1] - kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0) - kv_indices = self.cuda_graph_kv_indices - create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices, - seq_lens, - kv_indptr, - None, - kv_indices, - self.req_to_token.stride(0), - ) - kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs] - max_q_len = num_tokens_per_bs - - if self.use_mla and _use_mla_ps_kernel: - num_kv_splits = self.max_split_per_batch - - self.make_mla_meta_data( - qo_indptr, - kv_indptr, - kv_last_page_len, - self.work_metadata, - self.work_info_set, - self.work_indptr, - self.reduce_indptr, - self.reduce_final_map, - self.reduce_partial_map, - max_q_len, - fast_mode=fast_mode, - max_split_per_batch=num_kv_splits, - intra_batch_mode=intra_batch_mode, - ) - - work_metadata = self.work_metadata - work_info_set = self.work_info_set - work_indptr = self.work_indptr - - reduce_indptr = self.reduce_indptr - reduce_final_map = self.reduce_final_map - reduce_partial_map = self.reduce_partial_map - - self.forward_metadata = ForwardMetadata( - kv_indptr, - kv_indices, - qo_indptr, - kv_last_page_len, - max_q_len, - max_kv_len, - work_metadata=work_metadata, - work_info_set=work_info_set, - work_indptr=work_indptr, - reduce_indptr=reduce_indptr, - reduce_final_map=reduce_final_map, - reduce_partial_map=reduce_partial_map, - num_kv_splits=num_kv_splits, - ) - elif forward_mode.is_draft_extend(): - # EAGLE V1: Uses speculative_num_steps + 1 - num_tokens_per_bs = self.speculative_num_steps + 1 - qo_indptr = self.qo_indptr[: bs + 1] - qo_indptr[: bs + 1] = torch.arange( - 0, - bs * num_tokens_per_bs + 1, - step=num_tokens_per_bs, - dtype=torch.int32, - device=self.device, - ) - kv_indptr = self.kv_indptr[: bs + 1] - kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0) - kv_indices = self.cuda_graph_kv_indices - create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices, - seq_lens, - kv_indptr, - None, - kv_indices, - self.req_to_token.stride(0), - ) - - if self.use_mla: - kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs] - max_q_len = num_tokens_per_bs - - if _use_mla_ps_kernel: - num_kv_splits = self.max_split_per_batch - - self.make_mla_meta_data( - qo_indptr, - kv_indptr, - kv_last_page_len, - self.work_metadata, - self.work_info_set, - self.work_indptr, - self.reduce_indptr, - self.reduce_final_map, - self.reduce_partial_map, - max_q_len, - fast_mode=fast_mode, - max_split_per_batch=num_kv_splits, - intra_batch_mode=intra_batch_mode, - ) - - work_metadata = self.work_metadata - work_info_set = self.work_info_set - work_indptr = self.work_indptr - - reduce_indptr = self.reduce_indptr - reduce_final_map = self.reduce_final_map - reduce_partial_map = self.reduce_partial_map - - self.forward_metadata = ForwardMetadata( - kv_indptr, - kv_indices, - qo_indptr, - kv_last_page_len, - max_q_len, - max_kv_len, - work_metadata=work_metadata, - work_info_set=work_info_set, - work_indptr=work_indptr, - reduce_indptr=reduce_indptr, - reduce_final_map=reduce_final_map, - reduce_partial_map=reduce_partial_map, - num_kv_splits=num_kv_splits, - ) - else: - # Non-MLA draft_extend cuda graph: use triton extend kernel - self.forward_metadata = ForwardMetadata( - kv_indptr, - kv_indices, - qo_indptr, - None, - num_tokens_per_bs, - None, - custom_mask=None, - mask_indptr=None, - max_extend_len=num_tokens_per_bs, - ) - else: - raise ValueError(f"Invalid mode: {forward_mode=}") + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=None, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=seq_lens.cpu(), + ) def init_forward_metadata_replay_cuda_graph( self, @@ -1934,7 +1527,11 @@ class AiterAttnBackend(AttentionBackend): reduce_partial_map = None swa_page_table = None - max_kv_len = seq_lens_cpu.max().item() + max_kv_len = ( + seq_lens_cpu.max().item() + if seq_lens_cpu is not None + else torch.max(seq_lens).item() + ) if forward_mode.is_decode_or_idle(): qo_indptr = None diff --git a/python/sglang/srt/layers/attention/cutlass_mla_backend.py b/python/sglang/srt/layers/attention/cutlass_mla_backend.py index 05641ea1f..22781284d 100644 --- a/python/sglang/srt/layers/attention/cutlass_mla_backend.py +++ b/python/sglang/srt/layers/attention/cutlass_mla_backend.py @@ -153,24 +153,22 @@ class CutlassMLABackend(FlashInferMLAAttnBackend): forward_mode: ForwardMode, spec_info: Optional[SpecInput], ): - if forward_mode.is_decode_or_idle(): - if spec_info is None: - max_seqlen_pad = self.cuda_graph_kv_indices.shape[1] - - create_flashmla_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices, - seq_lens, - None, - self.cuda_graph_kv_indices, - self.req_to_token.stride(0), - self.cuda_graph_kv_indices.stride(0), - PAGED_SIZE=PAGE_SIZE, - ) - self.forward_metadata = CutlassMLADecodeMetadata( - self.cuda_graph_mla_workspace, - self.cuda_graph_kv_indices[:bs, :max_seqlen_pad], - ) + if forward_mode.is_decode_or_idle() and spec_info is None: + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=None, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=None, + ) + max_seqlen_pad = self.cuda_graph_kv_indices.shape[1] + self.forward_metadata = CutlassMLADecodeMetadata( + self.cuda_graph_mla_workspace, + self.cuda_graph_kv_indices[:bs, :max_seqlen_pad], + ) else: super().init_forward_metadata_capture_cuda_graph( bs, @@ -193,15 +191,11 @@ class CutlassMLABackend(FlashInferMLAAttnBackend): spec_info: Optional[SpecInput], seq_lens_cpu: Optional[torch.Tensor], ): - if forward_mode.is_decode_or_idle(): - assert seq_lens_cpu is not None - seq_lens = seq_lens[:bs] - create_flashmla_kv_indices_triton[(bs,)]( self.req_to_token, req_pool_indices[:bs], - seq_lens, + seq_lens[:bs], None, self.cuda_graph_kv_indices, self.req_to_token.stride(0), diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index e11c2d3a6..ac8a815be 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -749,48 +749,40 @@ class DeepseekV4AttnBackend( forward_mode: ForwardMode, spec_info: Optional[SpecInput], ) -> None: + from types import SimpleNamespace + assert req_pool_indices.size(0) == bs assert seq_lens.size(0) == bs bucket = _GraphBucket.of(forward_mode) - raw_type: Optional[type] = None if bucket == _GraphBucket.DECODE_OR_IDLE: - metadata = self.init_forward_metadata_decode( - max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - out_cache_loc=torch.zeros_like(seq_lens), - ) - raw_type = DSV4RawDecodeMetadata + dummy_cache_loc = torch.zeros_like(seq_lens) elif bucket == _GraphBucket.TARGET_VERIFY: - out_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs) - metadata = self.init_forward_metadata_target_verify( - max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - out_cache_loc=out_cache_loc, - use_prefill_cuda_graph=True, - ) - raw_type = DSV4RawVerifyMetadata - elif bucket == _GraphBucket.DRAFT_EXTEND: - num_tokens_per_bs = num_tokens // bs - metadata = self.init_forward_metadata_draft_extend( - max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_cpu=seq_lens.tolist(), - num_tokens_per_bs=num_tokens_per_bs, - use_prefill_cuda_graph=True, - ) + dummy_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs) else: - raise NotImplementedError(f"{forward_mode=} not supported yet") + dummy_cache_loc = None - self.cuda_graph_metadata_of_bucket_and_bs[bucket][bs] = metadata - self.forward_metadata = metadata - if raw_type is not None: - self._current_capture_raw = ( - metadata if isinstance(metadata, raw_type) else None - ) + self._replay_forward_batch = SimpleNamespace( + out_cache_loc=dummy_cache_loc, + forward_mode=forward_mode, + ) + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=int(seq_lens.sum().item()), + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=seq_lens.cpu(), + ) + # Preserve _current_capture_raw for on_after_cuda_graph_warmup + metadata = self.forward_metadata + self._current_capture_raw = ( + metadata + if isinstance(metadata, (DSV4RawDecodeMetadata, DSV4RawVerifyMetadata)) + else None + ) def init_forward_metadata_replay_cuda_graph( self, @@ -892,6 +884,11 @@ class DeepseekV4AttnBackend( ], bucket: _GraphBucket, ) -> None: + if bs not in self.cuda_graph_metadata_of_bucket_and_bs[bucket]: + # First call (from capture): store the new metadata directly. + self.cuda_graph_metadata_of_bucket_and_bs[bucket][bs] = temp_metadata + self.forward_metadata = temp_metadata + return chosen_metadata = self.cuda_graph_metadata_of_bucket_and_bs[bucket][bs] chosen_metadata.copy_(temp_metadata) self.forward_metadata = chosen_metadata diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py index 88f27563a..f400764d5 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py @@ -748,48 +748,40 @@ class DeepseekV4HipRadixBackend( forward_mode: ForwardMode, spec_info: Optional[SpecInput], ) -> None: + from types import SimpleNamespace + assert req_pool_indices.size(0) == bs assert seq_lens.size(0) == bs bucket = _GraphBucket.of(forward_mode) - raw_type: Optional[type] = None if bucket == _GraphBucket.DECODE_OR_IDLE: - metadata = self.init_forward_metadata_decode( - max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - out_cache_loc=torch.zeros_like(seq_lens), - ) - raw_type = DSV4RawDecodeMetadata + dummy_cache_loc = torch.zeros_like(seq_lens) elif bucket == _GraphBucket.TARGET_VERIFY: - out_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs) - metadata = self.init_forward_metadata_target_verify( - max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - out_cache_loc=out_cache_loc, - use_prefill_cuda_graph=True, - ) - raw_type = DSV4RawVerifyMetadata - elif bucket == _GraphBucket.DRAFT_EXTEND: - num_tokens_per_bs = num_tokens // bs - metadata = self.init_forward_metadata_draft_extend( - max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_cpu=seq_lens.tolist(), - num_tokens_per_bs=num_tokens_per_bs, - use_prefill_cuda_graph=True, - ) + dummy_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs) else: - raise NotImplementedError(f"{forward_mode=} not supported yet") + dummy_cache_loc = None - self.cuda_graph_metadata_of_bucket_and_bs[bucket][bs] = metadata - self.forward_metadata = metadata - if raw_type is not None: - self._current_capture_raw = ( - metadata if isinstance(metadata, raw_type) else None - ) + self._replay_forward_batch = SimpleNamespace( + out_cache_loc=dummy_cache_loc, + forward_mode=forward_mode, + ) + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=int(seq_lens.sum().item()), + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=seq_lens.cpu(), + ) + # Preserve _current_capture_raw for on_after_cuda_graph_warmup + metadata = self.forward_metadata + self._current_capture_raw = ( + metadata + if isinstance(metadata, (DSV4RawDecodeMetadata, DSV4RawVerifyMetadata)) + else None + ) def init_forward_metadata_replay_cuda_graph( self, @@ -891,6 +883,11 @@ class DeepseekV4HipRadixBackend( ], bucket: _GraphBucket, ) -> None: + if bs not in self.cuda_graph_metadata_of_bucket_and_bs[bucket]: + # First call (from capture): store the new metadata directly. + self.cuda_graph_metadata_of_bucket_and_bs[bucket][bs] = temp_metadata + self.forward_metadata = temp_metadata + return chosen_metadata = self.cuda_graph_metadata_of_bucket_and_bs[bucket][bs] chosen_metadata.copy_(temp_metadata) self.forward_metadata = chosen_metadata diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index 92479f3dc..11b305516 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -813,19 +813,21 @@ class DeepseekSparseAttnBackend( ), } - def init_forward_metadata_capture_cuda_graph( + def _build_forward_metadata_cuda_graph( self, bs: int, num_tokens: int, req_pool_indices: torch.Tensor, seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], + seq_lens_cpu: Optional[torch.Tensor], forward_mode: ForwardMode, spec_info: Optional[SpecInput], + out_cache_loc: Optional[torch.Tensor] = None, + actual_forward_mode: Optional["ForwardMode"] = None, ): + """Create and store DSAMetadata for a new batch size during CUDA graph capture.""" self.set_dsa_prefill_impl(forward_batch=None) - """Initialize forward metadata for capturing CUDA graph.""" if forward_mode.is_decode_or_idle(): # Normal Decode # Get sequence information @@ -847,11 +849,11 @@ class DeepseekSparseAttnBackend( ) seqlens_expanded = cache_seqlens_int32 - dsa_extend_seq_lens_list = [1] * num_tokens + dsa_extend_seq_lens_list = [1] * bs if self.dsa_decode_impl == "flashmla_kv": flashmla_metadata = self.decode_cuda_graph_metadata[ "flashmla_metadata" - ].slice(slice(0, num_tokens + 1)) + ].slice(slice(0, bs + 1)) flashmla_metadata.copy_( self._compute_flashmla_metadata( cache_seqlens=dsa_cache_seqlens_int32, @@ -969,6 +971,28 @@ class DeepseekSparseAttnBackend( self.decode_cuda_graph_metadata[bs] = metadata self.forward_metadata = metadata + def init_forward_metadata_capture_cuda_graph( + self, + bs: int, + num_tokens: int, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + encoder_lens: Optional[torch.Tensor], + forward_mode: ForwardMode, + spec_info: Optional[SpecInput], + ): + """Initialize forward metadata for capturing CUDA graph.""" + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=None, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=seq_lens.cpu(), + ) + def init_forward_metadata_replay_cuda_graph( self, bs: int, @@ -985,6 +1009,20 @@ class DeepseekSparseAttnBackend( """Initialize forward metadata for replaying CUDA graph.""" assert seq_lens_cpu is not None + if bs not in self.decode_cuda_graph_metadata: + self._build_forward_metadata_cuda_graph( + bs, + None, + req_pool_indices, + seq_lens, + seq_lens_cpu, + forward_mode, + spec_info, + out_cache_loc, + actual_forward_mode, + ) + return + self.set_dsa_prefill_impl(forward_batch=None) seq_lens = seq_lens[:bs] diff --git a/python/sglang/srt/layers/attention/dual_chunk_flashattention_backend.py b/python/sglang/srt/layers/attention/dual_chunk_flashattention_backend.py index fa0da5c46..00dc6e169 100644 --- a/python/sglang/srt/layers/attention/dual_chunk_flashattention_backend.py +++ b/python/sglang/srt/layers/attention/dual_chunk_flashattention_backend.py @@ -532,16 +532,13 @@ class DualChunkFlashAttentionBackend(AttentionBackend): ), } - def init_forward_metadata_capture_cuda_graph( + def _bind_metadata_buffers( self, bs: int, - num_tokens: int, req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], forward_mode: ForwardMode, - spec_info: Optional[None], ): + """Allocate persistent metadata buffers for CUDA graph capture.""" metadata = DualChunkFlashAttentionMetadata() if forward_mode.is_decode_or_idle(): @@ -580,6 +577,36 @@ class DualChunkFlashAttentionBackend(AttentionBackend): self.forward_metadata = metadata + def init_forward_metadata_capture_cuda_graph( + self, + bs: int, + num_tokens: int, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + encoder_lens: Optional[torch.Tensor], + forward_mode: ForwardMode, + spec_info: Optional[None], + ): + self._bind_metadata_buffers(bs, req_pool_indices, forward_mode) + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=None, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=seq_lens.cpu(), + ) + # Restore max_seq_len scalars — replay sets actual values but CUDA graph + # needs the safe upper bound baked in at capture time. + if forward_mode.is_decode_or_idle(): + md = self.forward_metadata + md.max_seq_len = self.max_context_len + md.max_seq_len_intra = self.max_context_len + md.max_seq_len_succ = self.max_context_len + md.max_seq_len_inter = self.max_context_len + def init_forward_metadata_replay_cuda_graph( self, bs: int, diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index 1d0d1b8ba..fe1f8de78 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -1700,43 +1700,37 @@ class FlashAttentionBackend(AttentionBackend): # For decoder-only models, skip encoder_metadata allocation self.encoder_metadata = {} - def init_forward_metadata_capture_cuda_graph( + def _bind_metadata_buffers( self, bs: int, num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, encoder_lens: Optional[torch.Tensor], forward_mode: ForwardMode, spec_info: Optional[SpecInput], - ): - """Initialize forward metadata for capturing CUDA graph.""" - metadata = FlashAttentionMetadata() + device: torch.device, + ) -> tuple: + """Create FlashAttentionMetadata with pre-allocated buffer slice refs. - # metadata_expand is needed for Spec Decoding when top k > 1 + Assigns all buffer slice references but does NOT fill data values. + Stores the new metadata object(s) in the appropriate lookup dicts. + Returns (metadata, metadata_expand). + """ + metadata = FlashAttentionMetadata() metadata_expand = FlashAttentionMetadata() - device = seq_lens.device if forward_mode.is_decode_or_idle(): if spec_info is not None: - # Draft Decode if self.topk <= 1: - # When topk = 1, we use the normal decode metadata + # Draft Decode topk=1 metadata.cache_seqlens_int32 = self.decode_cuda_graph_metadata[ "cache_seqlens" ][:bs] - metadata.max_seq_len_k = seq_lens.max().item() + ( - self.speculative_step_id + 1 - ) metadata.cu_seqlens_q = self.decode_cuda_graph_metadata[ "cu_seqlens_q" ][: bs + 1] - metadata.cu_seqlens_k = torch.nn.functional.pad( - torch.cumsum( - metadata.cache_seqlens_int32, dim=0, dtype=torch.int32 - ), - (1, 0), - ) + metadata.cu_seqlens_k = self.decode_cuda_graph_metadata[ + "cu_seqlens_k" + ][: bs + 1] metadata.page_table = self.decode_cuda_graph_metadata[ "page_table_draft_decode" ][:bs, :] @@ -1746,13 +1740,11 @@ class FlashAttentionBackend(AttentionBackend): ][:bs, :] self.decode_cuda_graph_metadata[bs] = metadata else: - # When top k > 1, we need two specific draft decode metadata, and then merge states - # 1. The first half of metadata for prefix tokens + # Draft Decode topk>1: two metadata objects metadata.cache_seqlens_int32 = ( self.draft_decode_metadata_topk_normal["cache_seqlens"][:bs] ) metadata.max_seq_len_q = self.topk - metadata.max_seq_len_k = seq_lens.max().item() metadata.cu_seqlens_q = self.draft_decode_metadata_topk_normal[ "cu_seqlens_q" ][: bs + 1] @@ -1763,7 +1755,6 @@ class FlashAttentionBackend(AttentionBackend): "page_table" ][:bs, :] - # 2. The second half of metadata for draft tokens (per_batch_num_tokens = topk) metadata_expand.cache_seqlens_int32 = ( self.draft_decode_metadata_topk_expand["cache_seqlens"][ : bs * self.topk @@ -1787,16 +1778,15 @@ class FlashAttentionBackend(AttentionBackend): self.draft_decode_metadata_topk_expand[bs] = metadata_expand else: # Normal Decode - # Get sequence information - metadata.cache_seqlens_int32 = seq_lens.to(torch.int32) - batch_size = len(seq_lens) - device = seq_lens.device - metadata.cu_seqlens_k = torch.nn.functional.pad( - torch.cumsum(seq_lens, dim=0, dtype=torch.int32), (1, 0) - ) - # Precompute maximum sequence length - metadata.max_seq_len_k = seq_lens.max().item() - # Precompute page table + metadata.cache_seqlens_int32 = self.decode_cuda_graph_metadata[ + "cache_seqlens" + ][:bs] + metadata.cu_seqlens_q = self.decode_cuda_graph_metadata["cu_seqlens_q"][ + : bs + 1 + ] + metadata.cu_seqlens_k = self.decode_cuda_graph_metadata["cu_seqlens_k"][ + : bs + 1 + ] metadata.page_table = self.decode_cuda_graph_metadata["page_table"][ :bs, : ] @@ -1804,70 +1794,32 @@ class FlashAttentionBackend(AttentionBackend): metadata.swa_page_table = self.decode_cuda_graph_metadata[ "swa_page_table" ][:bs, :] - # Precompute cumulative sequence lengths - metadata.cu_seqlens_q = torch.arange( - 0, batch_size + 1, dtype=torch.int32, device=device - ) self.decode_cuda_graph_metadata[bs] = metadata - self._maybe_update_local_attn_metadata_for_capture(metadata, batch_size) - - # Compute scheduler_metadata into pre-allocated buffer for CUDA graph capture - if self._sched_meta_buf is not None: - sched = self._compute_scheduler_metadata( - batch_size, - max(metadata.max_seq_len_k, 1), - metadata.cache_seqlens_int32, - metadata.cu_seqlens_q, - ) - if sched is not None: - n = sched.shape[0] - self._sched_meta_buf[:n] = sched - self._sched_meta_buf[n:] = 0 - metadata.scheduler_metadata = self._sched_meta_buf[:n] - elif forward_mode.is_target_verify(): if self.topk <= 1: metadata.cache_seqlens_int32 = self.target_verify_metadata[ "cache_seqlens" ][:bs] - metadata.cache_seqlens_int32.copy_( - (seq_lens + self.speculative_num_draft_tokens) - ) - metadata.max_seq_len_q = self.speculative_num_draft_tokens - metadata.max_seq_len_k = ( - seq_lens.max().item() + self.speculative_num_draft_tokens - ) - - metadata.cu_seqlens_q = torch.arange( - 0, - bs * self.speculative_num_draft_tokens + 1, - self.speculative_num_draft_tokens, - dtype=torch.int32, - device=device, - ) - + metadata.cu_seqlens_q = self.target_verify_metadata["cu_seqlens_q"][ + : bs + 1 + ] metadata.cu_seqlens_k = self.target_verify_metadata["cu_seqlens_k"][ : (bs + 1) ] - metadata.page_table = self.target_verify_metadata["page_table"][:bs, :] - if self.use_sliding_window_kv_pool: metadata.swa_page_table = self.target_verify_metadata[ "swa_page_table" ][:bs, :] - self.target_verify_metadata[bs] = metadata else: - # When topk > 1, we need two specific target verify metadata, and then merge states - # 1. The first half of metadata for prefix tokens + # Target Verify topk>1: two (or three with SWA) metadata objects metadata.cache_seqlens_int32 = self.target_verify_metadata_topk_normal[ "cache_seqlens" ][:bs] metadata.max_seq_len_q = self.speculative_num_draft_tokens - # metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item(), do this in replay metadata.cu_seqlens_q = self.target_verify_metadata_topk_normal[ "cu_seqlens_q" ][: bs + 1] @@ -1878,7 +1830,6 @@ class FlashAttentionBackend(AttentionBackend): "page_table" ][:bs, :] - # 2. The second half of metadata for draft tokens (per_batch_num_tokens = topk) metadata_expand.cache_seqlens_int32 = ( self.target_verify_metadata_topk_expand["cache_seqlens"][ : bs * self.speculative_num_draft_tokens @@ -1891,7 +1842,6 @@ class FlashAttentionBackend(AttentionBackend): metadata_expand.cu_seqlens_k = self.target_verify_metadata_topk_expand[ "cu_seqlens_k" ][: bs * self.speculative_num_draft_tokens + 1] - metadata_expand.page_table = self.target_verify_metadata_topk_expand[ "page_table" ][: bs * self.speculative_num_draft_tokens] @@ -1913,7 +1863,6 @@ class FlashAttentionBackend(AttentionBackend): metadata_swa.cu_seqlens_k = self.target_verify_metadata_topk_swa[ "cu_seqlens_k" ][: bs * self.speculative_num_draft_tokens + 1] - metadata_swa.page_table = self.target_verify_metadata_topk_swa[ "page_table" ][: bs * self.speculative_num_draft_tokens] @@ -1921,33 +1870,20 @@ class FlashAttentionBackend(AttentionBackend): metadata.swa_spec_metadata = metadata_swa elif forward_mode.is_draft_extend(include_v2=True): + num_tokens_per_bs = num_tokens // bs metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][ :bs ] - metadata.cache_seqlens_int32.copy_(seq_lens) - - num_tokens_per_bs = num_tokens // bs metadata.max_seq_len_q = num_tokens_per_bs - metadata.max_seq_len_k = seq_lens.max().item() - - metadata.cu_seqlens_q = torch.arange( - 0, - bs * num_tokens_per_bs + 1, - num_tokens_per_bs, - dtype=torch.int32, - device=device, - ) - + metadata.cu_seqlens_q = self.draft_extend_metadata["cu_seqlens_q"][: bs + 1] metadata.cu_seqlens_k = self.draft_extend_metadata["cu_seqlens_k"][ : (bs + 1) ] metadata.page_table = self.draft_extend_metadata["page_table"][:bs, :] - if self.use_sliding_window_kv_pool: metadata.swa_page_table = self.draft_extend_metadata["swa_page_table"][ :bs, : ] - self.draft_extend_metadata[bs] = metadata if encoder_lens is not None: @@ -1958,13 +1894,81 @@ class FlashAttentionBackend(AttentionBackend): metadata.encoder_cu_seqlens_k = self.encoder_metadata[ "encoder_cu_seqlens_k" ][: (encoder_bs + 1)] - metadata.encoder_page_table = self.encoder_metadata["encoder_page_table"][ :bs, : ] - self.forward_metadata = metadata - self.forward_metadata_spec_decode_expand = metadata_expand + return metadata, metadata_expand + + def init_forward_metadata_capture_cuda_graph( + self, + bs: int, + num_tokens: int, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + encoder_lens: Optional[torch.Tensor], + forward_mode: ForwardMode, + spec_info: Optional[SpecInput], + ): + """Initialize forward metadata for capturing CUDA graph.""" + seq_lens_cpu = seq_lens.cpu() + self._bind_metadata_buffers( + bs, num_tokens, encoder_lens, forward_mode, spec_info, seq_lens.device + ) + + if forward_mode.is_decode_or_idle() and spec_info is not None and self.topk > 1: + # topk>1 draft decode: replay needs out_cache_loc which capture doesn't have; + # set forward_metadata directly and let actual CUDA graph replay fill data. + self.forward_metadata = self.draft_decode_metadata_topk_normal[bs] + self.forward_metadata_spec_decode_expand = ( + self.draft_decode_metadata_topk_expand[bs] + ) + return + + if forward_mode.is_target_verify() and self.topk > 1: + # topk>1 target verify: replay needs spec_info.positions and .custom_mask + # which are not populated at capture time. + self.forward_metadata = self.target_verify_metadata_topk_normal[bs] + self.forward_metadata_spec_decode_expand = ( + self.target_verify_metadata_topk_expand[bs] + ) + return + + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=None, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=seq_lens_cpu, + ) + + if forward_mode.is_decode_or_idle() and spec_info is None: + # Local attention and scheduler metadata require capture-time slice sizing. + # Both depend on data already filled by replay above. + metadata = self.decode_cuda_graph_metadata[bs] + self._maybe_update_local_attn_metadata_for_capture(metadata, bs) + if self._sched_meta_buf is not None: + sched = self._compute_scheduler_metadata( + bs, + max(metadata.max_seq_len_k, 1), + metadata.cache_seqlens_int32, + metadata.cu_seqlens_q, + ) + if sched is not None: + n = sched.shape[0] + self._sched_meta_buf[:n] = sched + self._sched_meta_buf[n:] = 0 + metadata.scheduler_metadata = self._sched_meta_buf[:n] + + if forward_mode.is_draft_extend(include_v2=True): + # CUDA graph bakes max_seq_len_q as a constant. replay() sets it to + # max(num_accept_tokens_cpu) which is None/empty at capture time, + # falling back to 1. Restore the correct upper bound so the kernel + # sees num_tokens_per_bs (not 1) for all replays of this graph. + self.forward_metadata.max_seq_len_q = num_tokens // bs def init_forward_metadata_replay_cuda_graph( self, diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index ed28fd019..37ddc3e7c 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -557,6 +557,81 @@ class FlashInferAttnBackend(AttentionBackend): self.cuda_graph_qk_indptr = [x.clone() for x in self.kv_indptr] self.cuda_graph_qo_indptr = [x.clone() for x in self.kv_indptr] + def _create_decode_wrappers(self, bs: int, num_tokens: int) -> list: + return [ + BatchDecodeWithPagedKVCacheWrapper( + self.workspace_buffer, + "NHD", + backend=self.decode_backend, + use_cuda_graph=True, + use_tensor_cores=self.decode_use_tensor_cores, + paged_kv_indptr_buffer=self.kv_indptr[i][: num_tokens + 1], + paged_kv_indices_buffer=self.cuda_graph_kv_indices[i], + paged_kv_last_page_len_buffer=self.kv_last_page_len[:num_tokens], + ) + for i in range(self.num_wrappers) + ] + + def _create_prefill_wrappers(self, bs: int, use_custom_mask: bool = False) -> list: + # FlashInfer's prefill wrapper decides mask mode based on whether + # `custom_mask_buf` is initialized (not whether a custom mask is provided). + # For cases like DFLASH draft (ENCODER_ONLY / non-causal) we do NOT use a + # custom mask, so we must avoid initializing `custom_mask_buf`, otherwise + # FlashInfer will treat the (zero) buffer as a real mask and block attention. + wrappers = [] + for i in range(self.num_wrappers): + extra = ( + { + "custom_mask_buf": self.cuda_graph_custom_mask, + "mask_indptr_buf": self.cuda_graph_qk_indptr[i][: bs + 1], + } + if use_custom_mask + else {} + ) + wrappers.append( + BatchPrefillWithPagedKVCacheWrapper( + self.workspace_buffer, + "NHD", + use_cuda_graph=True, + backend=self.prefill_backend, + 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], + **extra, + ) + ) + return wrappers + + def _prepare_cuda_graph_metadata( + self, + bs: int, + num_tokens: int, + forward_mode: ForwardMode, + spec_info: Optional[SpecInput], + ) -> None: + if forward_mode.is_decode_or_idle(): + decode_wrappers = self._create_decode_wrappers(bs, num_tokens) + self.decode_cuda_graph_metadata[bs] = decode_wrappers + self.forward_metadata = DecodeMetadata(decode_wrappers) + elif ( + forward_mode.is_target_verify() + or forward_mode.is_draft_extend() + or forward_mode.is_dllm_extend() + ): + use_custom_mask = ( + forward_mode.is_target_verify() + and spec_info is not None + and getattr(spec_info, "custom_mask", None) is not None + ) + prefill_wrappers = self._create_prefill_wrappers(bs, use_custom_mask) + self.prefill_cuda_graph_metadata[bs] = prefill_wrappers + self.forward_metadata = PrefillMetadata( + prefill_wrappers, forward_mode.is_dllm_extend(), False + ) + else: + raise ValueError(f"Invalid mode: {forward_mode=}") + def init_forward_metadata_capture_cuda_graph( self, bs: int, @@ -567,148 +642,24 @@ class FlashInferAttnBackend(AttentionBackend): forward_mode: ForwardMode, spec_info: Optional[SpecInput], ): + seq_lens_sum = seq_lens.sum().item() + seq_lens_cpu = seq_lens.cpu() + self._prepare_cuda_graph_metadata(bs, num_tokens, forward_mode, spec_info) + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=seq_lens_sum, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=seq_lens_cpu, + ) + # fast_decode_plan requires _cached_module set by the initial full + # begin_forward call above; install it only after that first plan runs. if forward_mode.is_decode_or_idle(): - decode_wrappers = [] - for i in range(self.num_wrappers): - decode_wrappers.append( - BatchDecodeWithPagedKVCacheWrapper( - self.workspace_buffer, - "NHD", - backend=self.decode_backend, - use_cuda_graph=True, - use_tensor_cores=self.decode_use_tensor_cores, - paged_kv_indptr_buffer=self.kv_indptr[i][: num_tokens + 1], - paged_kv_indices_buffer=self.cuda_graph_kv_indices[i], - paged_kv_last_page_len_buffer=self.kv_last_page_len[ - :num_tokens - ], - ) - ) - seq_lens_sum = seq_lens.sum().item() - self.indices_updater_decode.update( - req_pool_indices, - seq_lens, - seq_lens.cpu(), # may add a little overhead in capture stage - seq_lens_sum, - decode_wrappers=decode_wrappers, - encoder_lens=encoder_lens, - spec_info=spec_info, - fixed_split_size=None, - disable_split_kv=self.disable_cuda_graph_kv_split, - ) - self.decode_cuda_graph_metadata[bs] = decode_wrappers - self.forward_metadata = DecodeMetadata(decode_wrappers) - for i in range(self.num_wrappers): - decode_wrappers[i].begin_forward = partial( - fast_decode_plan, decode_wrappers[i] - ) - elif forward_mode.is_target_verify(): - # FlashInfer's prefill wrapper decides mask mode based on whether - # `custom_mask_buf` is initialized (not whether a custom mask is provided). - # For cases like DFLASH draft (ENCODER_ONLY / non-causal) we do NOT use a - # custom mask, so we must avoid initializing `custom_mask_buf`, otherwise - # FlashInfer will treat the (zero) buffer as a real mask and block attention. - use_custom_mask = ( - spec_info is not None - and getattr(spec_info, "custom_mask", None) is not None - ) - prefill_wrappers = [] - for i in range(self.num_wrappers): - wrapper_kwargs = {} - if use_custom_mask: - wrapper_kwargs = { - "custom_mask_buf": self.cuda_graph_custom_mask, - "mask_indptr_buf": self.cuda_graph_qk_indptr[i][: bs + 1], - } - - prefill_wrappers.append( - BatchPrefillWithPagedKVCacheWrapper( - self.workspace_buffer, - "NHD", - use_cuda_graph=True, - backend=self.prefill_backend, - 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], - **wrapper_kwargs, - ) - ) - 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=None, - prefill_wrappers=prefill_wrappers, - use_ragged=False, - encoder_lens=encoder_lens, - spec_info=spec_info, - ) - self.prefill_cuda_graph_metadata[bs] = prefill_wrappers - self.forward_metadata = PrefillMetadata(prefill_wrappers, False, False) - elif forward_mode.is_draft_extend(): - prefill_wrappers = [] - for i in range(self.num_wrappers): - prefill_wrappers.append( - BatchPrefillWithPagedKVCacheWrapper( - self.workspace_buffer, - "NHD", - backend=self.prefill_backend, - 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=None, - prefill_wrappers=prefill_wrappers, - use_ragged=False, - encoder_lens=encoder_lens, - spec_info=spec_info, - ) - 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=self.prefill_backend, - 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=not self.use_paged, - 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=}") + for w in self.decode_cuda_graph_metadata[bs]: + w.begin_forward = partial(fast_decode_plan, w) def init_forward_metadata_replay_cuda_graph( self, @@ -733,19 +684,7 @@ class FlashInferAttnBackend(AttentionBackend): fixed_split_size=None, disable_split_kv=self.disable_cuda_graph_kv_split, ) - elif forward_mode.is_target_verify(): - 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=None, - prefill_wrappers=self.prefill_cuda_graph_metadata[bs], - use_ragged=False, - encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None, - spec_info=spec_info, - ) - elif forward_mode.is_draft_extend(): + elif forward_mode.is_target_verify() or forward_mode.is_draft_extend(): self.indices_updater_prefill.update( req_pool_indices[:bs], seq_lens[:bs], diff --git a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py index 61b6c49a5..cda57acb4 100644 --- a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py @@ -384,7 +384,13 @@ class FlashInferMLAAttnBackend(AttentionBackend): forward_mode: ForwardMode, spec_info: Optional[SpecInput], ): + seq_lens_sum = seq_lens.sum().item() + seq_lens_cpu = seq_lens.cpu() + if forward_mode.is_decode_or_idle(): + # Decode: create wrapper, run the initial full begin_forward (False), + # then install the fast plan. After that, call replay so the + # data-update path (update(True)) is also exercised during capture. decode_wrapper = BatchMLAPagedAttentionWrapper( self.workspace_buffer, use_cuda_graph=True, @@ -394,8 +400,6 @@ class FlashInferMLAAttnBackend(AttentionBackend): kv_len_arr=self.cuda_graph_kv_lens[:num_tokens], backend="auto", ) - - seq_lens_sum = seq_lens.sum().item() self.indices_updater_decode.update( req_pool_indices, seq_lens, @@ -406,9 +410,12 @@ class FlashInferMLAAttnBackend(AttentionBackend): ) self.decode_cuda_graph_metadata[bs] = decode_wrapper self.forward_metadata = DecodeMetadata(decode_wrapper) + # fast_mla_decode_plan requires _cached_module set by the initial + # begin_forward above; install it only after that call completes. decode_wrapper.plan = partial(fast_mla_decode_plan, decode_wrapper) - elif forward_mode.is_target_verify(): - verify_wrapper = BatchMLAPagedAttentionWrapper( + elif forward_mode.is_target_verify() or forward_mode.is_draft_extend(): + # Prefill: create wrapper and store — replay handles the update call. + prefill_wrapper = BatchMLAPagedAttentionWrapper( self.workspace_buffer, use_cuda_graph=True, qo_indptr=self.cuda_graph_qo_indptr[: bs + 1], @@ -417,43 +424,22 @@ class FlashInferMLAAttnBackend(AttentionBackend): kv_len_arr=self.cuda_graph_kv_lens[:bs], backend="auto", ) - seq_lens_sum = seq_lens.sum().item() - self.indices_updater_prefill.update( - req_pool_indices, - seq_lens, - seq_lens_sum, - prefix_lens=None, - prefill_wrapper_paged=verify_wrapper, - use_ragged=False, - spec_info=spec_info, - ) - self.prefill_cuda_graph_metadata[bs] = verify_wrapper - self.forward_metadata = PrefillMetadata(verify_wrapper, False) - elif forward_mode.is_draft_extend(): - draft_extend_wrapper = BatchMLAPagedAttentionWrapper( - self.workspace_buffer, - use_cuda_graph=True, - qo_indptr=self.cuda_graph_qo_indptr[: bs + 1], - kv_indptr=self.cuda_graph_kv_indptr[: bs + 1], - kv_indices=self.cuda_graph_kv_indices, - kv_len_arr=self.cuda_graph_kv_lens[:bs], - backend="auto", - ) - seq_lens_sum = seq_lens.sum().item() - self.indices_updater_prefill.update( - req_pool_indices, - seq_lens, - seq_lens_sum, - prefix_lens=None, - prefill_wrapper_paged=draft_extend_wrapper, - use_ragged=False, - spec_info=spec_info, - ) - self.prefill_cuda_graph_metadata[bs] = draft_extend_wrapper - self.forward_metadata = PrefillMetadata(draft_extend_wrapper, False) + self.prefill_cuda_graph_metadata[bs] = prefill_wrapper + self.forward_metadata = PrefillMetadata(prefill_wrapper, False) else: raise ValueError(f"Invalid mode: {forward_mode=}") + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=seq_lens_sum, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=seq_lens_cpu, + ) + def init_forward_metadata_replay_cuda_graph( self, bs: int, @@ -488,17 +474,7 @@ class FlashInferMLAAttnBackend(AttentionBackend): spec_info=spec_info, **self.fast_decode_kwargs, ) - elif forward_mode.is_target_verify(): - self.indices_updater_prefill.update( - req_pool_indices[:bs], - seq_lens[:bs], - seq_lens_sum, - prefix_lens=None, - prefill_wrapper_paged=self.prefill_cuda_graph_metadata[bs], - use_ragged=False, - spec_info=spec_info, - ) - elif forward_mode.is_draft_extend(): + elif forward_mode.is_target_verify() or forward_mode.is_draft_extend(): self.indices_updater_prefill.update( req_pool_indices[:bs], seq_lens[:bs], diff --git a/python/sglang/srt/layers/attention/flashmla_backend.py b/python/sglang/srt/layers/attention/flashmla_backend.py index c0bce60ce..ff77ed927 100644 --- a/python/sglang/srt/layers/attention/flashmla_backend.py +++ b/python/sglang/srt/layers/attention/flashmla_backend.py @@ -4,6 +4,7 @@ Support attention backend for FlashMLA. from __future__ import annotations +import logging from dataclasses import dataclass from typing import TYPE_CHECKING, Callable, Optional, Tuple, Union @@ -22,6 +23,7 @@ if TYPE_CHECKING: from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.speculative.spec_info import SpecInput +logger = logging.getLogger(__name__) PAGE_SIZE = 64 @@ -193,83 +195,16 @@ class FlashMLABackend(FlashInferMLAAttnBackend): forward_mode: ForwardMode, spec_info: Optional[SpecInput], ): - if forward_mode.is_decode_or_idle(): - max_seqlen_pad = triton.cdiv(seq_lens.max().item(), PAGE_SIZE) - - create_flashmla_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices, - seq_lens, - None, - self.cuda_graph_kv_indices, - self.req_to_token.stride(0), - self.cuda_graph_kv_indices.stride(0), - ) - num_q_heads = self.num_q_heads - - mla_metadata, num_splits = get_mla_metadata( - seq_lens.to(torch.int32), - num_q_heads, - 1, - is_fp8_kvcache=self.is_fp8_kvcache, - ) - - actual_num_sm_parts = mla_metadata.shape[0] - assert actual_num_sm_parts <= self.cuda_graph_mla_metadata.shape[0], ( - f"num_sm_parts {actual_num_sm_parts} exceeds preallocated max " - f"{self.cuda_graph_mla_metadata.shape[0]}" - ) - - self.cuda_graph_mla_metadata[:actual_num_sm_parts].copy_(mla_metadata) - self.cuda_graph_num_splits[: bs + 1].copy_(num_splits) - - self.cuda_graph_mla_metadata_view = self.cuda_graph_mla_metadata[ - :actual_num_sm_parts - ] - self.cuda_graph_num_splits_view = self.cuda_graph_num_splits[: bs + 1] - - self.forward_metadata = FlashMLADecodeMetadata( - self.cuda_graph_mla_metadata_view, - self.cuda_graph_num_splits_view, - self.cuda_graph_kv_indices[:bs, :max_seqlen_pad], - ) - - elif forward_mode.is_target_verify(): - seq_lens = seq_lens + self.num_draft_tokens - max_seqlen_pad = triton.cdiv(seq_lens.max().item(), PAGE_SIZE) - - create_flashmla_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices, - seq_lens, - None, - self.cuda_graph_kv_indices, - self.req_to_token.stride(0), - self.cuda_graph_kv_indices.stride(0), - ) - - mla_metadata, num_splits = get_mla_metadata( - seq_lens.to(torch.int32), - self.num_draft_tokens * self.num_q_heads, - 1, - is_fp8_kvcache=self.is_fp8_kvcache, - ) - - actual_num_sm_parts = mla_metadata.shape[0] - assert actual_num_sm_parts <= self.cuda_graph_mla_metadata.shape[0] - - self.cuda_graph_mla_metadata[:actual_num_sm_parts].copy_(mla_metadata) - self.cuda_graph_num_splits[: bs + 1].copy_(num_splits) - - self.cuda_graph_mla_metadata_view = self.cuda_graph_mla_metadata[ - :actual_num_sm_parts - ] - self.cuda_graph_num_splits_view = self.cuda_graph_num_splits[: bs + 1] - - self.forward_metadata = FlashMLADecodeMetadata( - self.cuda_graph_mla_metadata_view, - self.cuda_graph_num_splits_view, - self.cuda_graph_kv_indices[:bs, :max_seqlen_pad], + if forward_mode.is_decode_or_idle() or forward_mode.is_target_verify(): + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=None, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=None, ) else: super().init_forward_metadata_capture_cuda_graph( @@ -293,60 +228,21 @@ class FlashMLABackend(FlashInferMLAAttnBackend): spec_info: Optional[SpecInput], seq_lens_cpu: Optional[torch.Tensor], ): - if forward_mode.is_decode_or_idle(): - assert seq_lens_cpu is not None + if forward_mode.is_decode_or_idle() or forward_mode.is_target_verify(): seq_lens = seq_lens[:bs] - seq_lens_cpu = seq_lens_cpu[:bs] - max_seqlen_pad = triton.cdiv(seq_lens_cpu.max().item(), PAGE_SIZE) + seq_lens_cpu = seq_lens_cpu[:bs] if seq_lens_cpu is not None else None - create_flashmla_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices[:bs], - seq_lens, - None, - self.cuda_graph_kv_indices, - self.req_to_token.stride(0), - self.cuda_graph_kv_indices.stride(0), + if forward_mode.is_target_verify(): + seq_lens = seq_lens + self.num_draft_tokens + if seq_lens_cpu is not None: + seq_lens_cpu = seq_lens_cpu + self.num_draft_tokens + + seq_max = ( + seq_lens_cpu.max().item() + if seq_lens_cpu is not None + else seq_lens.max().item() ) - num_q_heads = self.num_q_heads - - mla_metadata, num_splits = get_mla_metadata( - seq_lens.to(torch.int32), - num_q_heads, - 1, - is_fp8_kvcache=self.is_fp8_kvcache, - ) - - actual_num_sm_parts = mla_metadata.shape[0] - - if actual_num_sm_parts != self.cuda_graph_mla_metadata_view.shape[0]: - import logging - - logger = logging.getLogger(__name__) - logger.warning( - f"num_sm_parts mismatch in CUDA Graph replay: " - f"capture={self.cuda_graph_mla_metadata_view.shape[0]}, " - f"replay={actual_num_sm_parts}. " - f"This may indicate batch size changed between capture and replay." - ) - self.cuda_graph_mla_metadata_view = self.cuda_graph_mla_metadata[ - :actual_num_sm_parts - ] - self.cuda_graph_num_splits_view = self.cuda_graph_num_splits[: bs + 1] - - self.cuda_graph_mla_metadata[:actual_num_sm_parts].copy_(mla_metadata) - self.cuda_graph_num_splits[: bs + 1].copy_(num_splits) - - self.forward_metadata.mla_metadata = self.cuda_graph_mla_metadata_view - self.forward_metadata.num_splits = self.cuda_graph_num_splits_view - self.forward_metadata.block_kv_indices = self.cuda_graph_kv_indices[ - :bs, :max_seqlen_pad - ] - - elif forward_mode.is_target_verify(): - seq_lens = seq_lens[:bs] + self.num_draft_tokens - seq_lens_cpu = seq_lens_cpu[:bs] + self.num_draft_tokens - max_seqlen_pad = triton.cdiv(seq_lens_cpu.max().item(), PAGE_SIZE) + max_seqlen_pad = triton.cdiv(seq_max, PAGE_SIZE) create_flashmla_kv_indices_triton[(bs,)]( self.req_to_token, @@ -358,29 +254,47 @@ class FlashMLABackend(FlashInferMLAAttnBackend): self.cuda_graph_kv_indices.stride(0), ) + q_head_mult = ( + self.num_draft_tokens if forward_mode.is_target_verify() else 1 + ) mla_metadata, num_splits = get_mla_metadata( seq_lens.to(torch.int32), - self.num_draft_tokens * self.num_q_heads, + q_head_mult * self.num_q_heads, 1, is_fp8_kvcache=self.is_fp8_kvcache, ) actual_num_sm_parts = mla_metadata.shape[0] + assert actual_num_sm_parts <= self.cuda_graph_mla_metadata.shape[0], ( + f"num_sm_parts {actual_num_sm_parts} exceeds preallocated max " + f"{self.cuda_graph_mla_metadata.shape[0]}" + ) - if actual_num_sm_parts != self.cuda_graph_mla_metadata_view.shape[0]: + if ( + self.cuda_graph_mla_metadata_view is None + or actual_num_sm_parts != self.cuda_graph_mla_metadata_view.shape[0] + ): + if self.cuda_graph_mla_metadata_view is not None: + logger.warning( + f"num_sm_parts mismatch in CUDA Graph replay: " + f"capture={self.cuda_graph_mla_metadata_view.shape[0]}, " + f"replay={actual_num_sm_parts}. " + f"This may indicate batch size changed between capture and replay." + ) self.cuda_graph_mla_metadata_view = self.cuda_graph_mla_metadata[ :actual_num_sm_parts ] - self.cuda_graph_num_splits_view = self.cuda_graph_num_splits[: bs + 1] + # num_splits has shape (bs+1,) — always update for the current bs. + self.cuda_graph_num_splits_view = self.cuda_graph_num_splits[: bs + 1] self.cuda_graph_mla_metadata[:actual_num_sm_parts].copy_(mla_metadata) self.cuda_graph_num_splits[: bs + 1].copy_(num_splits) - self.forward_metadata.mla_metadata = self.cuda_graph_mla_metadata_view - self.forward_metadata.num_splits = self.cuda_graph_num_splits_view - self.forward_metadata.block_kv_indices = self.cuda_graph_kv_indices[ - :bs, :max_seqlen_pad - ] + self.forward_metadata = FlashMLADecodeMetadata( + self.cuda_graph_mla_metadata_view, + self.cuda_graph_num_splits_view, + self.cuda_graph_kv_indices[:bs, :max_seqlen_pad], + ) else: super().init_forward_metadata_replay_cuda_graph( bs, diff --git a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index 04bbf870e..eaa78565d 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -403,8 +403,15 @@ class MambaAttnBackendBase(AttentionBackend): forward_mode: ForwardMode, spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], ): - self.forward_metadata = self._capture_metadata( - bs, req_pool_indices, forward_mode, spec_info + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=None, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=None, ) def init_forward_metadata_replay_cuda_graph( @@ -539,9 +546,12 @@ class MambaAttnBackendBase(AttentionBackend): spec_info: Optional[SpecInput], seq_lens_cpu: Optional[torch.Tensor], ): - num_padding = torch.count_nonzero( - seq_lens_cpu == self.get_cuda_graph_seq_len_fill_value() - ) + if seq_lens_cpu is None: + num_padding = 0 + else: + num_padding = torch.count_nonzero( + seq_lens_cpu == self.get_cuda_graph_seq_len_fill_value() + ) # Make sure forward metadata is correctly handled for padding reqs req_pool_indices[bs - num_padding :] = 0 mamba_indices = self.req_to_token_pool.get_mamba_indices(req_pool_indices) @@ -576,13 +586,17 @@ class MambaAttnBackendBase(AttentionBackend): # If topk > 1, we need to use retrieve_next_token and retrieve_next_sibling to handle the eagle tree custom attention mask if forward_mode.is_target_verify() and self.topk > 1: - bs_without_pad = spec_info.retrieve_next_token.shape[0] - self.retrieve_next_token_list[bs - 1][:bs_without_pad].copy_( - spec_info.retrieve_next_token - ) - self.retrieve_next_sibling_list[bs - 1][:bs_without_pad].copy_( - spec_info.retrieve_next_sibling - ) + if ( + spec_info is not None + and getattr(spec_info, "retrieve_next_token", None) is not None + ): + bs_without_pad = spec_info.retrieve_next_token.shape[0] + self.retrieve_next_token_list[bs - 1][:bs_without_pad].copy_( + spec_info.retrieve_next_token + ) + self.retrieve_next_sibling_list[bs - 1][:bs_without_pad].copy_( + spec_info.retrieve_next_sibling + ) return ForwardMetadata( query_start_loc=self.query_start_loc_list[bs - 1], mamba_cache_indices=self.state_indices_list[bs - 1], @@ -703,13 +717,15 @@ class Mamba2AttnBackend(MambaAttnBackendBase): forward_mode: ForwardMode, spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], ): - metadata = self._capture_metadata(bs, req_pool_indices, forward_mode, spec_info) - draft_token_num = spec_info.draft_token_num if spec_info is not None else 1 - self.forward_metadata = Mamba2Metadata.prepare_decode( - metadata, - seq_lens, - is_target_verify=forward_mode.is_target_verify(), - draft_token_num=draft_token_num, + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=None, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=None, ) def init_forward_metadata_replay_cuda_graph( diff --git a/python/sglang/srt/layers/attention/linear/lightning_backend.py b/python/sglang/srt/layers/attention/linear/lightning_backend.py index 3840c58aa..bc63db161 100644 --- a/python/sglang/srt/layers/attention/linear/lightning_backend.py +++ b/python/sglang/srt/layers/attention/linear/lightning_backend.py @@ -89,9 +89,15 @@ class LightningAttentionBackend(MambaAttnBackendBase): forward_mode: ForwardMode, spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], ): - metadata = self._capture_metadata(bs, req_pool_indices, forward_mode, spec_info) - self.forward_metadata = BailingLinearMetadata.prepare_decode( - metadata.query_start_loc, metadata.mamba_cache_indices, bs, seq_lens + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=None, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=None, ) def init_forward_metadata_replay_cuda_graph( diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index d08c15e00..5f78d61ca 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -283,11 +283,179 @@ class TritonAttnBackend(AttentionBackend): MAX_NUM_SEQ=SCHEDULE_SEQ, ) + def _fill_kv_indptr_and_indices( + self, + bs: int, + seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + kv_indices: torch.Tensor, + ) -> torch.Tensor: + kv_indptr = self.kv_indptr[: bs + 1] + kv_indptr[1:] = torch.cumsum(seq_lens, dim=0) + create_flashinfer_kv_indices_triton[(bs,)]( + self.req_to_token, + req_pool_indices, + seq_lens, + kv_indptr, + None, + kv_indices, + self.req_to_token.stride(0), + ) + return kv_indptr + + def _update_decode_kv_buffers( + self, + bs: int, + seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + ): + """Fill KV (and SWA) cuda-graph buffers for decode/idle mode. + + Returns ``(kv_indptr, window_kv_indptr, window_kv_lens)`` where + ``window_kv_lens`` is ``None`` when sliding-window is disabled. + """ + seq_lens = seq_lens[:bs] + req_pool_indices = req_pool_indices[:bs] + kv_indptr = self._fill_kv_indptr_and_indices( + bs, seq_lens, req_pool_indices, self.cuda_graph_kv_indices + ) + window_kv_indptr = self.window_kv_indptr + window_kv_lens = None + if self.sliding_window_size is not None and self.sliding_window_size > 0: + window_kv_indptr, _, window_kv_lens, _ = update_sliding_window_buffer( + self.window_kv_indptr, + self.req_to_token, + self.sliding_window_size, + seq_lens, + req_pool_indices, + bs, + token_to_kv_pool=self.token_to_kv_pool, + window_kv_indices=self.cuda_graph_window_kv_indices, + ) + return kv_indptr, window_kv_indptr, window_kv_lens + + def _update_target_verify_buffers( + self, + bs: int, + seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + spec_info, + ): + """Fill all cuda-graph buffers for target_verify mode. + + Returns the ForwardMetadata components: + ``(qo_indptr, kv_indptr, custom_mask, mask_indptr, + window_kv_indptr, window_kv_indices, window_num_kv_splits, window_kv_offsets)`` + """ + qo_indptr = self.qo_indptr[: bs + 1] + qo_indptr[: bs + 1] = torch.arange( + 0, + (1 + bs) * self.num_draft_tokens, + step=self.num_draft_tokens, + dtype=torch.int32, + device=self.device, + ) + kv_indptr = self._fill_kv_indptr_and_indices( + bs, seq_lens, req_pool_indices, self.cuda_graph_kv_indices + ) + window_kv_indptr = self.window_kv_indptr + window_kv_indices = None + window_num_kv_splits = None + window_kv_offsets = None + if self.sliding_window_size is not None and self.sliding_window_size > 0: + window_kv_indices = self.cuda_graph_window_kv_indices + window_num_kv_splits = self.cuda_graph_window_num_kv_splits + window_kv_offsets = self.cuda_graph_window_kv_offsets + window_kv_indptr, window_kv_indices, _, window_kv_offsets[:bs] = ( + update_sliding_window_buffer( + self.window_kv_indptr, + self.req_to_token, + self.sliding_window_size, + seq_lens[:bs], + req_pool_indices, + bs, + token_to_kv_pool=self.token_to_kv_pool, + window_kv_indices=window_kv_indices, + ) + ) + custom_mask = self.cuda_graph_custom_mask + if ( + spec_info is not None + and getattr(spec_info, "custom_mask", None) is not None + ): + custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.custom_mask + else: + custom_mask = None + seq_mask_len = self.num_draft_tokens * (seq_lens + self.num_draft_tokens) + mask_indptr = self.mask_indptr[: bs + 1] + mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0) + return ( + qo_indptr, + kv_indptr, + custom_mask, + mask_indptr, + window_kv_indptr, + window_kv_indices, + window_num_kv_splits, + window_kv_offsets, + ) + + def _update_draft_extend_buffers( + self, + bs: int, + seq_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + forward_mode: ForwardMode, + spec_info: Optional[SpecInput], + ): + """Fill QO + KV cuda-graph buffers for draft_extend mode. + + Returns ``(qo_indptr, kv_indptr, num_tokens_per_bs)``. + """ + seq_lens = seq_lens[:bs] + num_tokens_per_bs = self.speculative_num_steps + 1 + qo_indptr = self.qo_indptr[: bs + 1] + qo_indptr[: bs + 1] = torch.arange( + 0, + bs * num_tokens_per_bs + 1, + step=num_tokens_per_bs, + dtype=torch.int32, + device=self.device, + ) + if forward_mode.is_draft_extend_v2(): + # DRAFT_EXTEND_V2: seq_lens = prefix + extend (bumped by eagle_info_v2). + # Triton extend kernel receives extend K/V as separate tensors, so + # kv_indptr/kv_indices must cover only the prefix portion. + # extend_seq_lens_tensor is only attached to spec_info at real + # replay (eagle_draft_extend_cuda_graph_runner.replay); during the + # capture-time warmup it's absent, so fall back to zeros (matches + # the pre-unification capture path in #26651). Clamp at 0 because + # padded rows (raw_bs..bs) leave seq_lens at the fill value (1) + # while extend_seq_lens stays at num_tokens_per_bs, which would + # otherwise produce negative kv_lens; padded rows reference + # reserved req-pool slot 0 and their output is discarded. + if ( + spec_info is not None + and getattr(spec_info, "extend_seq_lens_tensor", None) is not None + ): + extend_seq_lens = spec_info.extend_seq_lens_tensor[:bs].to(torch.int32) + else: + extend_seq_lens = torch.zeros( + bs, dtype=torch.int32, device=seq_lens.device + ) + kv_lens = torch.clamp(seq_lens - extend_seq_lens, min=0).to(torch.int32) + else: + # DRAFT_EXTEND_V1: seq_lens = prefix only. + kv_lens = seq_lens + kv_indptr = self._fill_kv_indptr_and_indices( + bs, kv_lens, req_pool_indices, self.cuda_graph_kv_indices + ) + return qo_indptr, kv_indptr, num_tokens_per_bs + def init_forward_metadata(self, forward_batch: ForwardBatch): """Init auxiliary variables for triton attention backend.""" bs = forward_batch.batch_size - kv_indptr = self.kv_indptr window_kv_indptr = self.window_kv_indptr window_kv_indices = None window_num_kv_splits = None @@ -297,19 +465,14 @@ class TritonAttnBackend(AttentionBackend): if forward_batch.forward_mode.is_decode_or_idle(): if spec_info is None: - kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0) - kv_indptr = kv_indptr[: bs + 1] kv_indices = torch.empty( forward_batch.seq_lens_sum, dtype=torch.int64, device=self.device ) - create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - forward_batch.req_pool_indices, + kv_indptr = self._fill_kv_indptr_and_indices( + bs, forward_batch.seq_lens, - kv_indptr, - None, + forward_batch.req_pool_indices, kv_indices, - self.req_to_token.stride(0), ) # Sliding window if ( @@ -371,19 +534,14 @@ class TritonAttnBackend(AttentionBackend): device=self.device, ) # Different with flashinfer kv_indptr and kv_indices construction - kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0) - kv_indptr = kv_indptr[: bs + 1] kv_indices = torch.empty( - kv_indptr[-1], dtype=torch.int64, device=self.device + forward_batch.seq_lens_sum, dtype=torch.int64, device=self.device ) - create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - forward_batch.req_pool_indices, + kv_indptr = self._fill_kv_indptr_and_indices( + bs, forward_batch.seq_lens, - kv_indptr, - None, + forward_batch.req_pool_indices, kv_indices, - self.req_to_token.stride(0), ) if self.sliding_window_size is not None and self.sliding_window_size > 0: @@ -435,23 +593,16 @@ class TritonAttnBackend(AttentionBackend): attn_logits = None attn_lse = None else: - kv_indptr[1 : bs + 1] = torch.cumsum( - forward_batch.extend_prefix_lens, dim=0 - ) - kv_indptr = kv_indptr[: bs + 1] kv_indices = torch.empty( sum(forward_batch.extend_prefix_lens_cpu), dtype=torch.int64, device=self.device, ) - create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - forward_batch.req_pool_indices, + kv_indptr = self._fill_kv_indptr_and_indices( + bs, forward_batch.extend_prefix_lens, - kv_indptr, - None, + forward_batch.req_pool_indices, kv_indices, - self.req_to_token.stride(0), ) # Sliding window if self.sliding_window_size is not None and self.sliding_window_size > 0: @@ -578,6 +729,83 @@ class TritonAttnBackend(AttentionBackend): device=self.device, ) + def _build_cuda_graph_forward_metadata( + self, + bs: int, + forward_mode: ForwardMode, + spec_info: Optional[SpecInput], + ) -> ForwardMetadata: + """Construct ForwardMetadata from the current cuda-graph buffer state. + + Called by capture after the buffer-update helpers have already run + (either via replay or directly). All fields reference the same + ``self.cuda_graph_*`` tensors that the captured graph kernels will + read — the Python object is rebuilt each capture, but the underlying + GPU memory addresses are stable. + """ + swa = self.sliding_window_size is not None and self.sliding_window_size > 0 + if forward_mode.is_decode_or_idle(): + return ForwardMetadata( + attn_logits=self.cuda_graph_attn_logits, + attn_lse=self.cuda_graph_attn_lse, + max_extend_len=None, + num_kv_splits=self.cuda_graph_num_kv_splits, + kv_indptr=self.kv_indptr[: bs + 1], + kv_indices=self.cuda_graph_kv_indices, + qo_indptr=None, + custom_mask=None, + mask_indptr=None, + window_kv_indptr=self.window_kv_indptr[: bs + 1] if swa else None, + window_kv_indices=self.cuda_graph_window_kv_indices if swa else None, + window_num_kv_splits=( + self.cuda_graph_window_num_kv_splits if swa else None + ), + window_kv_offsets=None, + swa_attn_logits=self.cuda_graph_swa_attn_logits, + ) + elif forward_mode.is_target_verify(): + custom_mask = ( + self.cuda_graph_custom_mask + if spec_info is not None + and getattr(spec_info, "custom_mask", None) is not None + else None + ) + return ForwardMetadata( + attn_logits=None, + attn_lse=None, + max_extend_len=self.num_draft_tokens, + num_kv_splits=None, + kv_indptr=self.kv_indptr[: bs + 1], + kv_indices=self.cuda_graph_kv_indices, + qo_indptr=self.qo_indptr[: bs + 1], + custom_mask=custom_mask, + mask_indptr=self.mask_indptr[: bs + 1], + window_kv_indptr=self.window_kv_indptr[: bs + 1] if swa else None, + window_kv_indices=self.cuda_graph_window_kv_indices if swa else None, + window_num_kv_splits=( + self.cuda_graph_window_num_kv_splits if swa else None + ), + window_kv_offsets=self.cuda_graph_window_kv_offsets if swa else None, + ) + elif forward_mode.is_draft_extend(include_v2=True): + return ForwardMetadata( + attn_logits=None, + attn_lse=None, + max_extend_len=self.speculative_num_steps + 1, + num_kv_splits=None, + kv_indptr=self.kv_indptr[: bs + 1], + kv_indices=self.cuda_graph_kv_indices, + qo_indptr=self.qo_indptr[: bs + 1], + custom_mask=None, + mask_indptr=None, + window_kv_indptr=self.window_kv_indptr, + window_kv_indices=None, + window_num_kv_splits=None, + window_kv_offsets=None, + ) + else: + raise ValueError(f"Invalid forward mode: {forward_mode=} for CUDA Graph.") + def init_forward_metadata_capture_cuda_graph( self, bs: int, @@ -589,172 +817,42 @@ class TritonAttnBackend(AttentionBackend): spec_info: Optional[SpecInput], ): assert encoder_lens is None, "Not supported" - window_kv_indptr = self.window_kv_indptr - window_kv_indices = None - window_num_kv_splits = None - window_kv_offsets = None - swa_attn_logits = None - if forward_mode.is_decode_or_idle(): - if spec_info is None: - kv_indptr = self.kv_indptr - kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0) - kv_indptr = kv_indptr[: bs + 1] - kv_indices = self.cuda_graph_kv_indices - create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices, - seq_lens, - kv_indptr, - None, - kv_indices, - self.req_to_token.stride(0), - ) - if ( - self.sliding_window_size is not None - and self.sliding_window_size > 0 - ): - window_kv_indices = self.cuda_graph_window_kv_indices - window_num_kv_splits = self.cuda_graph_window_num_kv_splits - window_kv_indptr, window_kv_indices, _, _ = ( - update_sliding_window_buffer_cuda_graph( - self.window_kv_indptr, - window_kv_indices, - self.req_to_token, - self.sliding_window_size, - seq_lens[:bs], - req_pool_indices, - bs, - self.token_to_kv_pool, - ) - ) - else: - kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices + # Multi-step speculative decode: kv buffers come from spec_info rather + # than the cuda-graph pool, so replay is not involved for this path. + if forward_mode.is_decode_or_idle() and spec_info is not None: + self.forward_metadata = ForwardMetadata( + attn_logits=self.cuda_graph_attn_logits, + attn_lse=self.cuda_graph_attn_lse, + max_extend_len=None, + num_kv_splits=self.cuda_graph_num_kv_splits, + kv_indptr=spec_info.kv_indptr, + kv_indices=spec_info.kv_indices, + qo_indptr=None, + custom_mask=None, + mask_indptr=None, + window_kv_indptr=self.window_kv_indptr, + window_kv_indices=None, + window_num_kv_splits=None, + window_kv_offsets=None, + swa_attn_logits=self.cuda_graph_swa_attn_logits, + ) + return - attn_logits = self.cuda_graph_attn_logits - swa_attn_logits = self.cuda_graph_swa_attn_logits - attn_lse = self.cuda_graph_attn_lse - max_extend_len = None - num_kv_splits = self.cuda_graph_num_kv_splits - qo_indptr = None - custom_mask = None - mask_indptr = None - elif forward_mode.is_target_verify(): - qo_indptr = self.qo_indptr[: bs + 1] - qo_indptr[: bs + 1] = torch.arange( - 0, - (1 + bs) * self.num_draft_tokens, - step=self.num_draft_tokens, - dtype=torch.int32, - device=self.device, - ) - kv_indptr = self.kv_indptr[: bs + 1] - kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0) - kv_indices = self.cuda_graph_kv_indices - create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices, - seq_lens, - kv_indptr, - None, - kv_indices, - self.req_to_token.stride(0), - ) - - if self.sliding_window_size is not None and self.sliding_window_size > 0: - window_kv_indices = self.cuda_graph_window_kv_indices - window_num_kv_splits = self.cuda_graph_window_num_kv_splits - window_kv_offsets = self.cuda_graph_window_kv_offsets - window_kv_indptr, window_kv_indices, _, window_kv_offsets[:bs] = ( - update_sliding_window_buffer_cuda_graph( - self.window_kv_indptr, - window_kv_indices, - self.req_to_token, - self.sliding_window_size, - seq_lens[:bs], - req_pool_indices, - bs, - self.token_to_kv_pool, - ) - ) - - custom_mask = self.cuda_graph_custom_mask - if ( - spec_info is not None - and getattr(spec_info, "custom_mask", None) is not None - ): - custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.custom_mask - else: - custom_mask = None - seq_mask_len = self.num_draft_tokens * (seq_lens + self.num_draft_tokens) - mask_indptr = self.mask_indptr[: bs + 1] - mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0) - max_extend_len = self.num_draft_tokens - num_kv_splits = None - attn_logits = None - attn_lse = None - elif forward_mode.is_draft_extend(include_v2=True): - num_tokens_per_bs = self.speculative_num_steps + 1 - qo_indptr = self.qo_indptr[: bs + 1] - qo_indptr[: bs + 1] = torch.arange( - 0, - bs * num_tokens_per_bs + 1, - step=num_tokens_per_bs, - dtype=torch.int32, - device=self.device, - ) - kv_indptr = self.kv_indptr[: bs + 1] - if forward_mode.is_draft_extend_v2(): - # DRAFT_EXTEND_V2: seq_lens = prefix + extend (bumped by eagle_info_v2). - # Triton extend kernel receives extend K/V as separate tensors, so - # kv_indptr/kv_indices must cover only the prefix portion. - extend_seq_lens = ( - spec_info.extend_seq_lens_tensor[:bs].to(torch.int32) - if spec_info is not None - and getattr(spec_info, "extend_seq_lens_tensor", None) is not None - else torch.zeros(bs, dtype=torch.int32, device=self.device) - ) - kv_lens = (seq_lens - extend_seq_lens).to(torch.int32) - else: - # DRAFT_EXTEND_V1: seq_lens = prefix only. - kv_lens = seq_lens - kv_indptr[1 : bs + 1] = torch.cumsum(kv_lens, dim=0) - kv_indices = self.cuda_graph_kv_indices - create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices, - kv_lens, - kv_indptr, - None, - kv_indices, - self.req_to_token.stride(0), - ) - custom_mask = None - mask_indptr = None - max_extend_len = num_tokens_per_bs - num_kv_splits = None - attn_logits = None - attn_lse = None - else: - raise ValueError( - f"Invalid forward mode: {forward_mode=} for CUDA Graph capture." - ) - - self.forward_metadata = ForwardMetadata( - attn_logits, - attn_lse, - max_extend_len, - num_kv_splits, - kv_indptr, - kv_indices, - qo_indptr, - custom_mask, - mask_indptr, - window_kv_indptr, - window_kv_indices, - window_num_kv_splits, - window_kv_offsets, - swa_attn_logits=swa_attn_logits, + # Run the same buffer update as replay, then freeze the result into + # a ForwardMetadata whose tensor fields point into the cuda-graph buffers. + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=None, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=None, + ) + self.forward_metadata = self._build_cuda_graph_forward_metadata( + bs, forward_mode, spec_info ) def init_forward_metadata_replay_cuda_graph( @@ -770,138 +868,23 @@ class TritonAttnBackend(AttentionBackend): ): # NOTE: encoder_lens expected to be zeros or None if forward_mode.is_decode_or_idle(): - # Update kv_indptr, kv_indices - kv_indptr = self.kv_indptr - kv_indices = self.cuda_graph_kv_indices - num_kv_splits = self.cuda_graph_num_kv_splits - if spec_info is None: - kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens[:bs], dim=0) - kv_indptr = kv_indptr[: bs + 1] - create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices[:bs], - seq_lens[:bs], - kv_indptr, - None, - kv_indices, - self.req_to_token.stride(0), + assert spec_info is None, "Multi-step cuda graph init is not done here." + _, _, window_kv_lens = self._update_decode_kv_buffers( + bs, seq_lens, req_pool_indices + ) + self.get_num_kv_splits(self.cuda_graph_num_kv_splits[:bs], seq_lens[:bs]) + if window_kv_lens is not None: + self.get_num_kv_splits( + self.cuda_graph_window_num_kv_splits[:bs], window_kv_lens[:bs] ) - num_token = bs - if ( - self.sliding_window_size is not None - and self.sliding_window_size > 0 - ): - window_num_kv_splits = self.cuda_graph_window_num_kv_splits - window_kv_indices = self.cuda_graph_window_kv_indices - _, _, window_kv_lens, _ = update_sliding_window_buffer_cuda_graph( - self.window_kv_indptr, - window_kv_indices, - self.req_to_token, - self.sliding_window_size, - seq_lens[:bs], - req_pool_indices[:bs], - bs, - self.token_to_kv_pool, - ) - self.get_num_kv_splits( - window_num_kv_splits[:num_token], window_kv_lens[:bs] - ) - - else: - assert False, "Multi-step cuda graph init is not done here." - self.get_num_kv_splits(num_kv_splits[:num_token], seq_lens[:bs]) - elif forward_mode.is_target_verify(): - # Update qo_indptr, kv_indptr, kv_indices, custom_mask, mask_indptr bs = len(req_pool_indices) - qo_indptr = self.qo_indptr[: bs + 1] - qo_indptr[: bs + 1] = torch.arange( - 0, - (1 + bs) * self.num_draft_tokens, - step=self.num_draft_tokens, - dtype=torch.int32, - device=self.device, + self._update_target_verify_buffers( + bs, seq_lens, req_pool_indices, spec_info ) - kv_indptr = self.kv_indptr[: bs + 1] - kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0) - kv_indices = self.cuda_graph_kv_indices - create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices, - seq_lens, - kv_indptr, - None, - kv_indices, - self.req_to_token.stride(0), - ) - if self.sliding_window_size is not None and self.sliding_window_size > 0: - window_num_kv_splits = self.cuda_graph_window_num_kv_splits - window_kv_indices = self.cuda_graph_window_kv_indices - window_kv_offsets = self.cuda_graph_window_kv_offsets - _, _, window_kv_lens, window_kv_offsets[:bs] = ( - update_sliding_window_buffer_cuda_graph( - self.window_kv_indptr, - window_kv_indices, - self.req_to_token, - self.sliding_window_size, - seq_lens[:bs], - req_pool_indices, - bs, - self.token_to_kv_pool, - ) - ) - custom_mask = self.cuda_graph_custom_mask - if ( - spec_info is not None - and getattr(spec_info, "custom_mask", None) is not None - ): - custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.custom_mask - else: - custom_mask = None - seq_mask_len = self.num_draft_tokens * (seq_lens + self.num_draft_tokens) - mask_indptr = self.mask_indptr[: bs + 1] - mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0) elif forward_mode.is_draft_extend(include_v2=True): - seq_lens = seq_lens[:bs] - num_tokens_per_bs = self.speculative_num_steps + 1 - qo_indptr = self.qo_indptr[: bs + 1] - qo_indptr[: bs + 1] = torch.arange( - 0, - bs * num_tokens_per_bs + 1, - step=num_tokens_per_bs, - dtype=torch.int32, - device=self.device, - ) - kv_indptr = self.kv_indptr[: bs + 1] - if forward_mode.is_draft_extend_v2(): - # DRAFT_EXTEND_V2: seq_lens = prefix + extend (bumped by eagle_info_v2). - # Triton extend kernel receives extend K/V as separate tensors, so - # kv_indptr/kv_indices must cover only the prefix portion. - # Clamp at 0 because padded rows (raw_bs..bs) leave seq_lens at - # the fill value (1) while extend_seq_lens stays at num_tokens_per_bs, - # which would otherwise produce negative kv_lens; padded rows - # reference reserved req-pool slot 0 and their output is discarded. - assert ( - spec_info is not None - and getattr(spec_info, "extend_seq_lens_tensor", None) is not None - ), "DRAFT_EXTEND_V2 replay requires spec_info.extend_seq_lens_tensor" - kv_lens = torch.clamp( - seq_lens - spec_info.extend_seq_lens_tensor[:bs].to(torch.int32), - min=0, - ).to(torch.int32) - else: - # DRAFT_EXTEND_V1: seq_lens = prefix only. - kv_lens = seq_lens - kv_indptr[1 : bs + 1] = torch.cumsum(kv_lens, dim=0) - kv_indices = self.cuda_graph_kv_indices - create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices, - kv_lens, - kv_indptr, - None, - kv_indices, - self.req_to_token.stride(0), + self._update_draft_extend_buffers( + bs, seq_lens, req_pool_indices, forward_mode, spec_info ) else: raise ValueError( @@ -1490,18 +1473,26 @@ def update_sliding_window_buffer( seq_lens, req_pool_indices, bs, - device, + device=None, token_to_kv_pool=None, + window_kv_indices=None, ): + """Fill window KV buffers for sliding-window attention. + + Pass ``window_kv_indices`` to write into a pre-allocated buffer (CUDA-graph + path); omit it (or pass ``None``) to allocate a fresh tensor (eager path, + requires ``device``). + """ window_kv_lens = torch.minimum( seq_lens, torch.tensor(sliding_window_size), ) window_kv_indptr[1 : bs + 1] = torch.cumsum(window_kv_lens, dim=0) window_kv_indptr = window_kv_indptr[: bs + 1] - window_kv_indices = torch.empty( - window_kv_indptr[-1], dtype=torch.int64, device=device - ) + if window_kv_indices is None: + window_kv_indices = torch.empty( + window_kv_indptr[-1], dtype=torch.int64, device=device + ) window_kv_start_idx = seq_lens - window_kv_lens create_flashinfer_kv_indices_triton[(bs,)]( req_to_token, @@ -1512,47 +1503,6 @@ def update_sliding_window_buffer( window_kv_indices, req_to_token.stride(0), ) - # full to swa index mapping - if hasattr(token_to_kv_pool, "translate_loc_from_full_to_swa"): - kv_last_index = window_kv_indptr[-1] - # Flush before+after: window_kv_indices is a different tensor than out_cache_loc. - token_to_kv_pool.invalidate_loc_cache() - window_kv_indices[:kv_last_index] = ( - token_to_kv_pool.translate_loc_from_full_to_swa( - window_kv_indices[:kv_last_index] - ) - ) - token_to_kv_pool.invalidate_loc_cache() - return window_kv_indptr, window_kv_indices, window_kv_lens, window_kv_start_idx - - -def update_sliding_window_buffer_cuda_graph( - window_kv_indptr, - window_kv_indices, - req_to_token, - sliding_window_size, - seq_lens, - req_pool_indices, - bs, - token_to_kv_pool=None, -): - window_kv_lens = torch.minimum( - seq_lens, - torch.tensor(sliding_window_size), - ) - window_kv_indptr[1 : bs + 1] = torch.cumsum(window_kv_lens, dim=0) - window_kv_indptr = window_kv_indptr[: bs + 1] - window_kv_start_idx = seq_lens - window_kv_lens - create_flashinfer_kv_indices_triton[(bs,)]( - req_to_token, - req_pool_indices, - window_kv_lens, - window_kv_indptr, - window_kv_start_idx, - window_kv_indices, - req_to_token.stride(0), - ) - # full to swa index mapping if hasattr(token_to_kv_pool, "translate_loc_from_full_to_swa"): kv_last_index = window_kv_indptr[-1] # Flush before+after: window_kv_indices is a different tensor than out_cache_loc. diff --git a/python/sglang/srt/layers/attention/trtllm_mha_backend.py b/python/sglang/srt/layers/attention/trtllm_mha_backend.py index 068791b02..69c1409d2 100644 --- a/python/sglang/srt/layers/attention/trtllm_mha_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mha_backend.py @@ -303,42 +303,29 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): ), } - def init_forward_metadata_capture_cuda_graph( + def _build_cuda_graph_metadata( self, bs: int, num_tokens: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - encoder_lens: Optional[torch.Tensor], forward_mode: ForwardMode, - spec_info: Optional[SpecInput], - ): - """Initialize metadata for CUDA graph capture.""" + spec_info, + device: torch.device, + ) -> "TRTLLMMHAMetadata": + """Create TRTLLMMHAMetadata with pre-allocated buffer slice refs, stored in the dict.""" metadata = TRTLLMMHAMetadata() - device = seq_lens.device if forward_mode.is_decode_or_idle(): if spec_info is not None: - # Draft Decode - # Here we only support topk = 1 for now. + # Draft Decode (topk = 1) metadata.cache_seqlens_int32 = self.decode_cuda_graph_metadata[ "cache_seqlens" ][:bs] - metadata.cache_seqlens_int32.copy_( - seq_lens + self.speculative_step_id + 1 - ) - metadata.max_seq_len_k = seq_lens.max().item() + ( - self.speculative_step_id + 1 - ) metadata.cu_seqlens_q = self.decode_cuda_graph_metadata["cu_seqlens_q"][ : bs + 1 ] - metadata.cu_seqlens_k = torch.nn.functional.pad( - torch.cumsum( - metadata.cache_seqlens_int32, dim=0, dtype=torch.int32 - ), - (1, 0), - ) + metadata.cu_seqlens_k = self.decode_cuda_graph_metadata["cu_seqlens_k"][ + : bs + 1 + ] metadata.page_table = self.decode_cuda_graph_metadata[ "page_table_draft_decode" ][:bs, :] @@ -351,20 +338,15 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): self.decode_cuda_graph_metadata[bs] = metadata else: # Normal Decode - # Get sequence information - metadata.cache_seqlens_int32 = seq_lens[:bs].to(torch.int32) - batch_size = len(seq_lens) - metadata.cu_seqlens_k = torch.nn.functional.pad( - torch.cumsum(seq_lens, dim=0, dtype=torch.int32), (1, 0) - ) - - # Precompute maximum sequence length - metadata.max_seq_len_k = seq_lens.max().item() - # Precompute cumulative sequence lengths + metadata.cache_seqlens_int32 = self.decode_cuda_graph_metadata[ + "cache_seqlens" + ][:bs] metadata.cu_seqlens_q = torch.arange( - 0, batch_size + 1, dtype=torch.int32, device=device + 0, bs + 1, dtype=torch.int32, device=device + ) + metadata.cu_seqlens_k = torch.zeros( + bs + 1, dtype=torch.int32, device=device ) - # Precompute page table metadata.page_table = self.decode_cuda_graph_metadata["page_table"][ :bs, : ] @@ -376,29 +358,18 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): ) self.decode_cuda_graph_metadata[bs] = metadata elif forward_mode.is_target_verify(): - # Target Verify - # Here we only support topk = 1 for now. + # Target Verify (topk = 1) tokens_per_req = num_tokens // bs metadata.cache_seqlens_int32 = self.target_verify_metadata["cache_seqlens"][ :bs ] - metadata.cache_seqlens_int32.copy_(seq_lens + tokens_per_req) - - metadata.cu_seqlens_q = torch.arange( - 0, - bs * tokens_per_req + 1, - tokens_per_req, - dtype=torch.int32, - device=device, - ) - + metadata.cu_seqlens_q = self.target_verify_metadata["cu_seqlens_q"][ + : bs + 1 + ] metadata.cu_seqlens_k = self.target_verify_metadata["cu_seqlens_k"][ - : (bs + 1) + : bs + 1 ] - metadata.max_seq_len_q = tokens_per_req - metadata.max_seq_len_k = seq_lens.max().item() + tokens_per_req - metadata.page_table = self.target_verify_metadata["page_table"][:bs, :] self._bind_swa_page_table( metadata, @@ -406,29 +377,15 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): "swa_page_table", bs, ) - self.target_verify_metadata[bs] = metadata elif forward_mode.is_draft_extend(): + num_tokens_per_bs = num_tokens // bs metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][ :bs ] - metadata.cache_seqlens_int32.copy_(seq_lens) - num_tokens_per_bs = num_tokens // bs - metadata.cu_seqlens_q = torch.arange( - 0, - bs * num_tokens_per_bs + 1, - num_tokens_per_bs, - dtype=torch.int32, - device=device, - ) - - metadata.cu_seqlens_k = self.draft_extend_metadata["cu_seqlens_k"][ - : (bs + 1) - ] - num_tokens_per_bs = num_tokens // bs + metadata.cu_seqlens_q = self.draft_extend_metadata["cu_seqlens_q"][: bs + 1] + metadata.cu_seqlens_k = self.draft_extend_metadata["cu_seqlens_k"][: bs + 1] metadata.max_seq_len_q = num_tokens_per_bs - metadata.max_seq_len_k = seq_lens.max().item() - metadata.page_table = self.draft_extend_metadata["page_table"][:bs, :] self._bind_swa_page_table( metadata, @@ -436,9 +393,41 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): "swa_page_table", bs, ) - self.draft_extend_metadata[bs] = metadata - self.forward_metadata = metadata + + return metadata + + def init_forward_metadata_capture_cuda_graph( + self, + bs: int, + num_tokens: int, + req_pool_indices: torch.Tensor, + seq_lens: torch.Tensor, + encoder_lens: Optional[torch.Tensor], + forward_mode: ForwardMode, + spec_info: Optional[SpecInput], + ): + """Initialize metadata for CUDA graph capture.""" + seq_lens_cpu = seq_lens.cpu() + self._build_cuda_graph_metadata( + bs, num_tokens, forward_mode, spec_info, seq_lens.device + ) + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=None, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=seq_lens_cpu, + ) + if forward_mode.is_draft_extend(): + # CUDA graph bakes max_seq_len_q as a constant. replay() sets it to + # max(num_accept_tokens_cpu) which is None/empty at capture time, + # falling back to 1. Restore the correct upper bound so the kernel + # sees num_tokens_per_bs (not 1) for all replays of this graph. + self.forward_metadata.max_seq_len_q = num_tokens // bs def init_forward_metadata_replay_cuda_graph( self, diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index 493185647..932df4306 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -444,6 +444,44 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): super().init_cuda_graph_state(max_bs, max_num_tokens, kv_indices_buf) + def _init_cuda_graph_metadata( + self, + bs: int, + num_tokens: int, + forward_mode: ForwardMode, + seq_lens: torch.Tensor, + device: torch.device, + ): + """Allocate persistent metadata buffers for CUDA graph capture.""" + metadata = TRTLLMMLADecodeMetadata() + + if forward_mode.is_target_verify(): + metadata.seq_lens_k = torch.zeros((bs,), dtype=torch.int32, device=device) + elif forward_mode.is_draft_extend(include_v2=True): + num_tokens_per_bs = num_tokens // bs + metadata.max_seq_len_q = num_tokens_per_bs + metadata.sum_seq_lens_q = num_tokens_per_bs * bs + metadata.cu_seqlens_q = torch.arange( + 0, + bs * num_tokens_per_bs + 1, + num_tokens_per_bs, + dtype=torch.int32, + device=device, + ) + metadata.seq_lens_q = torch.full( + (bs,), num_tokens_per_bs, dtype=torch.int32, device=device + ) + metadata.seq_lens_k = torch.zeros((bs,), dtype=torch.int32, device=device) + + # Capture with full width so future longer sequences are safe during replay. + max_blocks_per_seq = self._calc_padded_blocks(self.max_context_len) + block_kv_indices = self.decode_cuda_graph_kv_indices[:bs, :max_blocks_per_seq] + metadata.block_kv_indices = block_kv_indices + metadata.max_seq_len_k = self.max_context_len + + self.decode_cuda_graph_metadata[bs] = metadata + self.forward_decode_metadata = metadata + def init_forward_metadata_capture_cuda_graph( self, bs: int, @@ -472,60 +510,19 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): spec_info, ) - metadata = TRTLLMMLADecodeMetadata() - - if forward_mode.is_target_verify(): - seq_lens = seq_lens + self.num_draft_tokens - metadata.seq_lens_k = torch.zeros( - (bs,), dtype=torch.int32, device=seq_lens.device - ) - metadata.seq_lens_k.copy_(seq_lens.to(dtype=torch.int32)) - elif forward_mode.is_draft_extend(include_v2=True): - num_tokens_per_bs = num_tokens // bs - metadata.max_seq_len_q = num_tokens_per_bs - metadata.sum_seq_lens_q = num_tokens_per_bs * bs - metadata.cu_seqlens_q = torch.arange( - 0, - bs * num_tokens_per_bs + 1, - num_tokens_per_bs, - dtype=torch.int32, - device=seq_lens.device, - ) - metadata.seq_lens_q = torch.full( - (bs,), num_tokens_per_bs, dtype=torch.int32, device=seq_lens.device - ) - # NOTE(draft_extend seq_len handling): - # forward_batch.seq_lens is the seq_lens of the prev_context + verified tokens. - # To account for pad_draft_extend_query, we need seq_lens = prev_context + max_draft_tokens. - # This will ensure queries align with kvs correctly when calling - # flashinfer.decode.trtllm_batch_decode_with_kv_cache_mla. - seq_lens = seq_lens - metadata.seq_lens_q + metadata.max_seq_len_q - metadata.seq_lens_k = torch.zeros( - (bs,), dtype=torch.int32, device=seq_lens.device - ) - metadata.seq_lens_k.copy_(seq_lens.to(dtype=torch.int32)) - - # Custom fast-path for decode/idle. - # Capture with full width so future longer sequences are safe during replay - max_blocks_per_seq = self._calc_padded_blocks(self.max_context_len) - block_kv_indices = self.decode_cuda_graph_kv_indices[:bs, :max_blocks_per_seq] - - create_flashmla_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices, - seq_lens, - None, - block_kv_indices, - self.req_to_token.stride(0), - max_blocks_per_seq, - PAGED_SIZE=self.page_size, + self._init_cuda_graph_metadata( + bs, num_tokens, forward_mode, seq_lens, seq_lens.device + ) + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=None, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=seq_lens.cpu(), ) - - metadata.block_kv_indices = block_kv_indices - metadata.max_seq_len_k = self.max_context_len - - self.decode_cuda_graph_metadata[bs] = metadata - self.forward_decode_metadata = metadata def init_forward_metadata_replay_cuda_graph( self, diff --git a/python/sglang/srt/layers/attention/wave_backend.py b/python/sglang/srt/layers/attention/wave_backend.py index ebdadc9bf..ff304a13d 100644 --- a/python/sglang/srt/layers/attention/wave_backend.py +++ b/python/sglang/srt/layers/attention/wave_backend.py @@ -390,6 +390,39 @@ class WaveAttnBackend(AttentionBackend): device=self.device, ) + def _build_cuda_graph_forward_metadata( + self, + bs: int, + forward_mode: ForwardMode, + spec_info: Optional[SpecInput], + ) -> ForwardMetadata: + if forward_mode.is_decode_or_idle(): + return ForwardMetadata( + attn_logits=self.cuda_graph_attn_logits, + attn_lse=self.cuda_graph_attn_lse, + max_extend_len=None, + num_kv_splits=self.cuda_graph_num_kv_splits, + kv_indptr=self.kv_indptr[: bs + 1], + kv_indices=self.cuda_graph_kv_indices, + qo_indptr=None, + custom_mask=None, + mask_indptr=None, + ) + elif forward_mode.is_target_verify(): + return ForwardMetadata( + attn_logits=None, + attn_lse=None, + max_extend_len=self.num_draft_tokens, + num_kv_splits=None, + kv_indptr=self.kv_indptr[: bs + 1], + kv_indices=self.cuda_graph_kv_indices, + qo_indptr=self.qo_indptr[: bs + 1], + custom_mask=self.cuda_graph_custom_mask, + mask_indptr=self.mask_indptr[: bs + 1], + ) + else: + raise ValueError(f"Invalid forward mode: {forward_mode=} for CUDA Graph.") + def init_forward_metadata_capture_cuda_graph( self, bs: int, @@ -402,76 +435,34 @@ class WaveAttnBackend(AttentionBackend): ): assert encoder_lens is None, "Not supported" - if forward_mode.is_decode_or_idle(): - if spec_info is None: - kv_indptr = self.kv_indptr - kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0) - kv_indptr = kv_indptr[: bs + 1] - kv_indices = self.cuda_graph_kv_indices - create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices, - seq_lens, - kv_indptr, - None, - kv_indices, - self.req_to_token.stride(0), - ) - else: - kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices - - attn_logits = self.cuda_graph_attn_logits - attn_lse = self.cuda_graph_attn_lse - max_extend_len = None - num_kv_splits = self.cuda_graph_num_kv_splits - qo_indptr = None - custom_mask = None - mask_indptr = None - elif forward_mode.is_target_verify(): - qo_indptr = self.qo_indptr[: bs + 1] - qo_indptr[: bs + 1] = torch.arange( - 0, - (1 + bs) * self.num_draft_tokens, - step=self.num_draft_tokens, - dtype=torch.int32, - device=self.device, - ) - kv_indptr = self.kv_indptr[: bs + 1] - kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0) - kv_indices = self.cuda_graph_kv_indices - create_flashinfer_kv_indices_triton[(bs,)]( - self.req_to_token, - req_pool_indices, - seq_lens, - kv_indptr, - None, - kv_indices, - self.req_to_token.stride(0), + # Multi-step speculative decode: kv buffers come from spec_info rather than + # the cuda-graph pool, so replay is not involved for this path. + if forward_mode.is_decode_or_idle() and spec_info is not None: + self.forward_metadata = ForwardMetadata( + attn_logits=self.cuda_graph_attn_logits, + attn_lse=self.cuda_graph_attn_lse, + max_extend_len=None, + num_kv_splits=self.cuda_graph_num_kv_splits, + kv_indptr=spec_info.kv_indptr, + kv_indices=spec_info.kv_indices, + qo_indptr=None, + custom_mask=None, + mask_indptr=None, ) + return - custom_mask = self.cuda_graph_custom_mask - seq_mask_len = self.num_draft_tokens * (seq_lens + self.num_draft_tokens) - mask_indptr = self.mask_indptr[: bs + 1] - mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0) - max_extend_len = self.num_draft_tokens - num_kv_splits = None - attn_logits = None - attn_lse = None - else: - raise ValueError( - f"Invalid forward mode: {forward_mode=} for CUDA Graph capture." - ) - - self.forward_metadata = ForwardMetadata( - attn_logits, - attn_lse, - max_extend_len, - num_kv_splits, - kv_indptr, - kv_indices, - qo_indptr, - custom_mask, - mask_indptr, + self.init_forward_metadata_replay_cuda_graph( + bs=bs, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_sum=None, + encoder_lens=encoder_lens, + forward_mode=forward_mode, + spec_info=spec_info, + seq_lens_cpu=None, + ) + self.forward_metadata = self._build_cuda_graph_forward_metadata( + bs, forward_mode, spec_info ) def init_forward_metadata_replay_cuda_graph( @@ -485,9 +476,7 @@ class WaveAttnBackend(AttentionBackend): spec_info: Optional[SpecInput], seq_lens_cpu: Optional[torch.Tensor], ): - # NOTE: encoder_lens expected to be zeros or None if forward_mode.is_decode_or_idle(): - # Update kv_indptr, kv_indices kv_indptr = self.kv_indptr kv_indices = self.cuda_graph_kv_indices num_kv_splits = self.cuda_graph_num_kv_splits @@ -510,7 +499,6 @@ class WaveAttnBackend(AttentionBackend): num_token = spec_info.kv_indptr.shape[0] - 1 self.get_num_kv_splits(num_kv_splits[:num_token], seq_lens[:bs]) elif forward_mode.is_target_verify(): - # Update qo_indptr, kv_indptr, kv_indices, custom_mask, mask_indptr bs = len(req_pool_indices) qo_indptr = self.qo_indptr[: bs + 1] qo_indptr[: bs + 1] = torch.arange( diff --git a/test/registered/attention/unittests/dense/test_tbo.py b/test/registered/attention/unittests/dense/test_tbo.py index 48c4a80a3..a5867451f 100644 --- a/test/registered/attention/unittests/dense/test_tbo.py +++ b/test/registered/attention/unittests/dense/test_tbo.py @@ -7,6 +7,7 @@ import torch from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS from sglang.srt.layers.attention.tbo_backend import TboAttnBackend from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.srt.utils import get_device_sm from sglang.test.test_utils import CustomTestCase sys.path.insert(0, str(Path(__file__).resolve().parents[1])) @@ -21,6 +22,9 @@ from sglang.test.kits.attention_unittest.attention_methods.dense_attention impor replace_backend, run_dense_fixture_eager, ) +from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_runner import ( + _prepare_spec_verify_batch, +) register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200") register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large") @@ -53,11 +57,27 @@ class TestTboAttnDenseAttentionBackendCorrectness(CustomTestCase): extend_lens=(16,), ) + # Mirrors ``runner_fa3_eagle_verify_chain`` in test_fa3.py — the smallest + # case shape that drives a real TARGET_VERIFY CUDA-graph capture through + # FlashAttention's per-bs metadata dicts. + TARGET_VERIFY_CAPTURE_CASE = DenseAttentionCase( + name="tbo_fa3_target_verify_chain_capture", + backend="fa3", + forward_mode=ForwardMode.TARGET_VERIFY, + num_heads=4, + num_kv_heads=4, + page_size=16, + prefix_lens=(4, 7), + extend_lens=(3, 3), + ) + def _build_and_wrap(self, case: DenseAttentionCase): fixture = build_dense_attention_fixture(self, case) try: - primary = ATTENTION_BACKENDS["triton"](fixture.runner) - children = [ATTENTION_BACKENDS["triton"](fixture.runner) for _ in range(2)] + primary = ATTENTION_BACKENDS[case.backend](fixture.runner) + children = [ + ATTENTION_BACKENDS[case.backend](fixture.runner) for _ in range(2) + ] except (AssertionError, ImportError, ModuleNotFoundError) as exc: self.skipTest(f"tbo child backend unavailable: {exc}") wrapper = TboAttnBackend(primary=primary, children=children) @@ -69,6 +89,56 @@ class TestTboAttnDenseAttentionBackendCorrectness(CustomTestCase): expected = expected_dense_fixture_output(fixture) torch.testing.assert_close(actual, expected, atol=DENSE_ATOL, rtol=DENSE_RTOL) + @unittest.skipIf( + get_device_sm() >= 100 or get_device_sm() < 80, + "FA3 backend requires SM 80-90", + ) + def test_tbo_target_verify_cuda_graph_capture_delegates_to_primary_capture(self): + """TBO capture must invoke ``primary.init_forward_metadata_capture_cuda_graph``, + not ``primary.init_forward_metadata_replay_cuda_graph``. + + Backends like FlashAttention store per-bs metadata in dicts populated + only by their capture path (via ``_bind_metadata_buffers``). If TBO + short-circuits its capture to its own replay (which delegates to + ``primary.replay``), those dicts are empty and replay raises + ``KeyError: bs``. Reproduces the deepep-4-gpu-h100 failure where + ``flashattention_backend.target_verify_metadata[bs]`` lookup blew up + during ``init_device_graphs``. + + Asserts capture completes without raising — numerical correctness of + the captured graph is covered by per-backend spec-verify tests. + """ + case = self.TARGET_VERIFY_CAPTURE_CASE + fixture = self._build_and_wrap(case) + wrapper = fixture.backend + batch = fixture.forward_batch + + # Wire TARGET_VERIFY batch state + EAGLE chain (topk=1) spec_info, + # mirroring what the per-backend spec-verify runner sets up. + _prepare_spec_verify_batch( + case, + batch, + topk=1, + spec_kind="eagle", + device=str(batch.seq_lens.device), + ) + + capture_bs = case.batch_size + num_tokens = sum(case.extend_lens) + wrapper.init_cuda_graph_state(max_bs=capture_bs, max_num_tokens=num_tokens) + # This is the failing call before the fix: TBO.capture delegating to + # primary.replay (instead of primary.capture) reads an unpopulated + # ``target_verify_metadata[bs]`` dict and raises KeyError. + wrapper.init_forward_metadata_capture_cuda_graph( + bs=capture_bs, + num_tokens=num_tokens, + req_pool_indices=batch.req_pool_indices, + seq_lens=batch.seq_lens, + encoder_lens=batch.encoder_lens, + forward_mode=batch.forward_mode, + spec_info=batch.spec_info, + ) + if __name__ == "__main__": unittest.main()