From a813224e7808ff17c4fca27a228e6b076ceb930c Mon Sep 17 00:00:00 2001 From: AMD-yanfeiwang Date: Wed, 16 Sep 2026 23:52:11 +0800 Subject: [PATCH] [ROCm][DSV4] Enable breakable CUDA graph prefill (#37810) Co-authored-by: Duyi-Wang --- .../attention/dsv4_attn_metadata_kernels.py | 9 +- .../srt/layers/attention/base_attn_backend.py | 4 + .../deepseek_v4_backend_hip_radix.py | 321 ++++++++++++-- .../managers/scheduler_components/dp_attn.py | 7 +- .../runner/prefill_cuda_graph_runner.py | 20 + .../amd/test_dsv4_hip_bcg_metadata.py | 419 ++++++++++++++++++ .../scheduler_components/test_dp_attn.py | 40 ++ .../runner/test_prefill_cuda_graph_padding.py | 15 +- 8 files changed, 778 insertions(+), 57 deletions(-) create mode 100644 test/registered/amd/test_dsv4_hip_bcg_metadata.py diff --git a/python/sglang/kernels/ops/attention/dsv4_attn_metadata_kernels.py b/python/sglang/kernels/ops/attention/dsv4_attn_metadata_kernels.py index bbb28952c..4070be229 100644 --- a/python/sglang/kernels/ops/attention/dsv4_attn_metadata_kernels.py +++ b/python/sglang/kernels/ops/attention/dsv4_attn_metadata_kernels.py @@ -142,10 +142,11 @@ def expand_prefill_causally( seq_lens_casual = torch.nn.functional.pad( seq_lens_casual, (0, pad_size), value=1 ) - req_pool_indices_repeated = torch.nn.functional.pad( - req_pool_indices_repeated, - (0, pad_size), - value=req_pool_indices_repeated[-1].item(), + req_pool_indices_repeated = torch.cat( + ( + req_pool_indices_repeated, + req_pool_indices_repeated[-1:].expand(pad_size), + ) ) return ExpandPrefillCausallyResult( seq_lens_casual=seq_lens_casual, diff --git a/python/sglang/srt/layers/attention/base_attn_backend.py b/python/sglang/srt/layers/attention/base_attn_backend.py index 5e2729fd1..6279ce7db 100644 --- a/python/sglang/srt/layers/attention/base_attn_backend.py +++ b/python/sglang/srt/layers/attention/base_attn_backend.py @@ -153,6 +153,10 @@ class AttentionBackend(ABC): # object during capture, and refresh its dynamic fields before each replay. use_captured_forward_metadata_for_breakable_cuda_graph: bool = False + # Backends may keep MIXED prefill eager under DP attention when replaying + # the EXTEND graph is a known serving-performance regression. + prefer_eager_mixed_prefill_under_dp_attention: bool = False + # True when prefill graph metadata can use ForwardBatch.max_seq_len_override. supports_prefill_cuda_graph_max_context_size: bool = False 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 be5f606fe..68e6d5f21 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 @@ -147,6 +147,31 @@ class UnifiedKvMetadata: assign_fields=[], ) + def refresh_for_breakable_cuda_graph_replay_( + self, other: UnifiedKvMetadata + ) -> None: + copy_metadata( + src=other, + dst=self, + check_eq_fields=[], + copy_fields=[ + "swa_loc", + "swa_indices", + "swa_indptr", + "hca_indices", + "hca_indptr", + "csa_indices", + "csa_indptr", + "pf_state_slot", + "pf_chunk_start", + "pf_cu_q", + "pf_final_pos", + "verify_store_state_slot", + "c4_out_loc", + "c128_out_loc", + ], + ) + @dataclass class DSV4AttnMetadata: @@ -237,6 +262,51 @@ class DSV4AttnMetadata: ], ) + def refresh_for_breakable_cuda_graph_replay_(self, other: DSV4AttnMetadata) -> None: + assert self.c4_sparse_topk == other.c4_sparse_topk + assert self.page_size == other.page_size + assert self.cuda_int32_kwargs == other.cuda_int32_kwargs + + tensor_copy_fields = [ + "raw_out_loc", + "seq_lens_casual", + "positions_casual", + "swa_out_cache_loc", + "c4_out_loc", + "c128_out_loc", + "page_table", + "swa_page_indices", + "swa_topk_lengths", + "c128_page_indices", + "c128_topk_lengths_clamp1", + "c128_topk_lengths_raw", + "c4_topk_lengths_raw", + "c4_topk_lengths_clamp1", + "c4_sparse_topk_lengths", + "c4_sparse_topk_lengths_raw", + "c4_sparse_page_indices", + "c4_sparse_raw_indices", + ] + for field_name in tensor_copy_fields: + src_val = getattr(other, field_name) + dst_val = getattr(self, field_name) + if src_val is None and dst_val is None: + continue + assert src_val is not None and dst_val is not None, ( + f"{field_name=} {src_val=} {dst_val=}" + ) + dst_val.copy_(src_val) + + if self.unified is None and other.unified is None: + pass + else: + assert self.unified is not None and other.unified is not None + self.unified.refresh_for_breakable_cuda_graph_replay_(other.unified) + + self.c0_flashmla_metadata = other.c0_flashmla_metadata + self.c4_flashmla_metadata = other.c4_flashmla_metadata + self.c128_flashmla_metadata = other.c128_flashmla_metadata + def init_compression_metadata(self, unified_swa_pages: int = 0): assert self.page_table.dim() == 2 assert self.raw_out_loc.shape == self.seq_lens_casual.shape, ( @@ -383,6 +453,37 @@ class DSV4Metadata: self.c128_compress_metadata, src=other.c128_compress_metadata ) + def refresh_for_breakable_cuda_graph_replay_(self, other: DSV4Metadata) -> None: + self.core_attn_metadata.refresh_for_breakable_cuda_graph_replay_( + other.core_attn_metadata + ) + maybe_copy_inplace(self.indexer_metadata, src=other.indexer_metadata) + maybe_copy_inplace(self.c4_compress_metadata, src=other.c4_compress_metadata) + maybe_copy_inplace( + self.c128_compress_metadata, src=other.c128_compress_metadata + ) + + if self.fp4_k_write_metadata is None and other.fp4_k_write_metadata is None: + pass + else: + assert ( + self.fp4_k_write_metadata is not None + and other.fp4_k_write_metadata is not None + ) + for captured, replay in zip( + self.fp4_k_write_metadata, + other.fp4_k_write_metadata, + strict=True, + ): + captured.copy_(replay) + + if self.fp4_q_positions is None and other.fp4_q_positions is None: + pass + else: + assert self.fp4_q_positions is not None + assert other.fp4_q_positions is not None + self.fp4_q_positions.copy_(other.fp4_q_positions) + @dataclass class DSV4RawVerifyMetadata: @@ -439,6 +540,9 @@ class DeepseekV4HipRadixBackend( # TboAttnBackend reads this to skip children in the *_graph paths only. tbo_supports_cuda_graph = False supports_ragged_verify_graph: bool = True + use_captured_forward_metadata_for_breakable_cuda_graph: bool = True + # MIXED BCG replay regresses ROCm DSV4 DP-attention serving throughput. + prefer_eager_mixed_prefill_under_dp_attention: bool = True def __init__( self, @@ -585,13 +689,22 @@ class DeepseekV4HipRadixBackend( need_compress=need_compress, is_prefill=True, ) + # Normal prefill starts with a conservative exact_num_tokens=False. + # Its CPU length mirror proves the exact query count without a D2H sync. + host_proves_exact_num_tokens = ( + need_compress + and not attach_decode_streams + and extend_seq_lens_cpu is not None + and sum(extend_seq_lens_cpu) == num_tokens + ) self._attach_unified_kv_prefill_meta( core_attn_metadata, req_pool_indices, + req_pool_indices_repeated, seq_lens, extend_seq_lens, num_tokens, - exact_num_tokens=exact_num_tokens, + exact_num_tokens=exact_num_tokens or host_proves_exact_num_tokens, ) if attach_decode_streams: # Target-verify runs through the unified_kv DECODE kernel, so build @@ -608,33 +721,41 @@ class DeepseekV4HipRadixBackend( ) if not need_compress: create = _create_dummy_paged_compress_data - elif compress_gpu_plan: - create = functools.partial( - create_paged_compressor_data, - is_prefill=True, - token_to_kv_pool=self.token_to_kv_pool, - req_to_token=self.req_to_token, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_cpu=None, - extend_lens=extend_seq_lens, - extend_lens_cpu=None, - num_q_tokens=num_tokens, - use_prefill_cuda_graph=use_prefill_cuda_graph, - ) else: - create = functools.partial( - create_paged_compressor_data, - is_prefill=True, - token_to_kv_pool=self.token_to_kv_pool, - req_to_token=self.req_to_token, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - seq_lens_cpu=seq_lens_cpu, - extend_lens=extend_seq_lens, - extend_lens_cpu=extend_seq_lens_cpu, - use_prefill_cuda_graph=use_prefill_cuda_graph, - ) + + def create(compress_ratio: Literal[4, 128]): + use_graph_plan = use_prefill_cuda_graph and not ( + compress_ratio == 128 and envs.SGLANG_OPT_USE_ONLINE_COMPRESS.get() + ) + if compress_gpu_plan or use_graph_plan: + return create_paged_compressor_data( + compress_ratio=compress_ratio, + is_prefill=True, + token_to_kv_pool=self.token_to_kv_pool, + req_to_token=self.req_to_token, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_cpu=None, + extend_lens=extend_seq_lens, + extend_lens_cpu=None, + num_q_tokens=( + out_cache_loc.shape[0] if use_graph_plan else num_tokens + ), + use_prefill_cuda_graph=use_prefill_cuda_graph, + ) + return create_paged_compressor_data( + compress_ratio=compress_ratio, + is_prefill=True, + token_to_kv_pool=self.token_to_kv_pool, + req_to_token=self.req_to_token, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + seq_lens_cpu=seq_lens_cpu, + extend_lens=extend_seq_lens, + extend_lens_cpu=extend_seq_lens_cpu, + use_prefill_cuda_graph=False, + ) + return DSV4Metadata( core_attn_metadata, indexer_metadata, @@ -1147,10 +1268,13 @@ class DeepseekV4HipRadixBackend( else None ) - def init_forward_metadata(self, forward_batch: ForwardBatch) -> None: - if self.mtp_enabled and forward_batch.forward_mode.is_idle(): - return - + def _build_forward_metadata( + self, + forward_batch: ForwardBatch, + *, + max_seq_len_override: Optional[int] = None, + use_prefill_cuda_graph: bool = False, + ): req_pool_indices = forward_batch.req_pool_indices seq_lens = forward_batch.seq_lens.to(torch.int32) seq_lens_cpu = forward_batch.seq_lens_cpu @@ -1158,7 +1282,11 @@ class DeepseekV4HipRadixBackend( assert self.swa_page_size % SWA_WINDOW == 0 and self.page_size % 128 == 0 assert seq_lens_cpu is not None - max_seq_len = int(seq_lens_cpu.max().item()) + max_seq_len = ( + max_seq_len_override + if max_seq_len_override is not None + else int(seq_lens_cpu.max().item()) + ) if forward_batch.forward_mode.is_decode_or_idle(): # DSv4 bakes this step's KV write target (c4/c128) into metadata, @@ -1211,16 +1339,61 @@ class DeepseekV4HipRadixBackend( num_tokens=sum(extend_seq_lens_cpu), extend_seq_lens=extend_seq_lens, extend_seq_lens_cpu=extend_seq_lens_cpu, + extend_start_loc=forward_batch.extend_start_loc, need_compress=not is_draft, + use_prefill_cuda_graph=use_prefill_cuda_graph, exact_num_tokens=is_draft, ) else: raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}") - self.forward_metadata = metadata + return metadata + + def init_forward_metadata(self, forward_batch: ForwardBatch) -> None: + if self.mtp_enabled and forward_batch.forward_mode.is_idle(): + return + + self.forward_metadata = self._build_forward_metadata(forward_batch) self.init_forward_metadata_in_graph(forward_batch) self._refresh_fp4_prefill_workspace(forward_batch) + def init_forward_metadata_for_breakable_cuda_graph_capture( + self, forward_batch: ForwardBatch + ): + self.forward_metadata = self._build_forward_metadata( + forward_batch, + max_seq_len_override=self.MAX_SEQ_LEN_FOR_CAPTURE, + use_prefill_cuda_graph=True, + ) + self.init_forward_metadata_in_graph(forward_batch) + self._refresh_fp4_prefill_workspace(forward_batch) + assert isinstance(self.forward_metadata, DSV4Metadata) + return self.forward_metadata + + def prepare_forward_metadata_for_breakable_cuda_graph_replay( + self, + capture_metadata, + forward_batch: ForwardBatch, + *, + static_forward_batch: Optional[ForwardBatch] = None, + ) -> None: + replay_batch = ( + static_forward_batch if static_forward_batch is not None else forward_batch + ) + replay_metadata = self._build_forward_metadata( + replay_batch, + max_seq_len_override=self.MAX_SEQ_LEN_FOR_CAPTURE, + use_prefill_cuda_graph=True, + ) + self.forward_metadata = replay_metadata + self.init_forward_metadata_in_graph(replay_batch) + + assert isinstance(capture_metadata, DSV4Metadata) + assert isinstance(replay_metadata, DSV4Metadata) + capture_metadata.refresh_for_breakable_cuda_graph_replay_(replay_metadata) + self.forward_metadata = capture_metadata + self._refresh_fp4_prefill_workspace(replay_batch) + def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int) -> None: self.cuda_graph_metadata_of_bucket_and_bs: Dict[ _GraphBucket, @@ -1327,6 +1500,7 @@ class DeepseekV4HipRadixBackend( self, core: DSV4AttnMetadata, req_pool_indices: torch.Tensor, + req_pool_indices_repeated: torch.Tensor, seq_lens: torch.Tensor, extend_seq_lens: torch.Tensor, num_tokens: int, @@ -1352,11 +1526,33 @@ class DeepseekV4HipRadixBackend( ) if core.unified is None: core.unified = UnifiedKvMetadata() - core.unified.pf_state_slot = req_pool_indices[bid] - core.unified.pf_chunk_start = (seq_lens - extend_seq_lens)[bid] + state_slot = req_pool_indices[bid] + chunk_start = (seq_lens - extend_seq_lens)[bid] cu_q_per_req = torch.cumsum(extend_seq_lens, dim=0) - extend_seq_lens - core.unified.pf_cu_q = cu_q_per_req[bid] - core.unified.pf_final_pos = (seq_lens - 1)[bid] + cu_q = cu_q_per_req[bid] + final_pos = (seq_lens - 1)[bid] + + padded_num_tokens = core.positions_casual.shape[0] + assert num_tokens <= padded_num_tokens + if num_tokens < padded_num_tokens: + pad_size = padded_num_tokens - num_tokens + state_slot = torch.cat( + (state_slot, req_pool_indices_repeated[num_tokens:padded_num_tokens]) + ) + chunk_start = F.pad(chunk_start, (0, pad_size), value=0) + cu_q = F.pad(cu_q, (0, pad_size), value=0) + # Padded positions are zero. final_pos=win makes the SWA store's + # `pos <= final_pos - win` guard skip every padded row. + final_pos = F.pad( + final_pos, + (0, pad_size), + value=self.token_to_kv_pool.unified_swa_window, + ) + + core.unified.pf_state_slot = state_slot + core.unified.pf_chunk_start = chunk_start + core.unified.pf_cu_q = cu_q + core.unified.pf_final_pos = final_pos def _forward_unified_kv( self, @@ -1739,16 +1935,12 @@ class DeepseekV4HipRadixBackend( swa_page_indices = core_attn_metadata.swa_page_indices swa_topk_lengths = core_attn_metadata.swa_topk_lengths - if self.mtp_enabled: - if swa_page_indices.shape[0] != q.shape[0]: - swa_page_indices = _pad_tensor_to_size( - swa_page_indices, q.shape[0], value=0 - ) - - if swa_topk_lengths.shape[0] != q.shape[0]: - swa_topk_lengths = _pad_tensor_to_size( - swa_topk_lengths, q.shape[0], value=1 - ) + swa_page_indices = _match_num_queries(swa_page_indices, q.shape[0], value=0) + swa_topk_lengths = _match_num_queries(swa_topk_lengths, q.shape[0], value=1) + extra_indices = _match_num_queries(extra_indices, q.shape[0], value=-1) + extra_topk_lengths = _match_num_queries( + extra_topk_lengths, q.shape[0], value=1 + ) if q.ndim == 3: q = q.unsqueeze(1) @@ -1992,6 +2184,33 @@ class DeepseekV4MultiStepBackend(DeepseekV4HipRadixBackend): for i in range(self.speculative_num_steps - 1): self.attn_backends[i].init_forward_metadata(forward_batch) + def init_forward_metadata_for_breakable_cuda_graph_capture( + self, forward_batch: ForwardBatch + ): + return [ + self.attn_backends[ + i + ].init_forward_metadata_for_breakable_cuda_graph_capture(forward_batch) + for i in range(self.speculative_num_steps - 1) + ] + + def prepare_forward_metadata_for_breakable_cuda_graph_replay( + self, + capture_metadata, + forward_batch: ForwardBatch, + *, + static_forward_batch: Optional[ForwardBatch] = None, + ) -> None: + assert len(capture_metadata) == self.speculative_num_steps - 1 + for i in range(self.speculative_num_steps - 1): + self.attn_backends[ + i + ].prepare_forward_metadata_for_breakable_cuda_graph_replay( + capture_metadata[i], + forward_batch, + static_forward_batch=static_forward_batch, + ) + def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): for i in range(self.speculative_num_steps): self.attn_backends[i].init_cuda_graph_state(max_bs, max_num_tokens) @@ -2001,7 +2220,13 @@ class DeepseekV4MultiStepBackend(DeepseekV4HipRadixBackend): backend.on_after_cuda_graph_warmup() -def _pad_tensor_to_size(tensor: torch.Tensor, size: int, *, value: int = 0): +def _match_num_queries( + tensor: Optional[torch.Tensor], size: int, *, value: int +) -> Optional[torch.Tensor]: + if tensor is None or tensor.shape[0] == size: + return tensor + if tensor.shape[0] > size: + return tensor[:size] if value == 0: return torch.cat( [tensor, tensor.new_zeros(size - tensor.shape[0], *tensor.shape[1:])], diff --git a/python/sglang/srt/managers/scheduler_components/dp_attn.py b/python/sglang/srt/managers/scheduler_components/dp_attn.py index dd2384936..7edb862a9 100644 --- a/python/sglang/srt/managers/scheduler_components/dp_attn.py +++ b/python/sglang/srt/managers/scheduler_components/dp_attn.py @@ -289,9 +289,9 @@ def _local_prefill_cuda_graph_vote( model_config, ) -> bool: """This rank's vote for the prefill graph (min-reduced across dp - ranks). Extend/mixed batches vote their own replayability; a decode - batch eligible for the decode->extend conversion votes as its 1-token- - extend view, so the vote and the post-sync conversion always agree.""" + ranks). Extend and mixed batches share the runner's rank-local replay + policy. A decode batch eligible for the decode->extend conversion votes as + its 1-token-extend view, so the vote and post-sync conversion agree.""" if local_batch is None or local_batch.forward_mode.is_idle(): return True if not coordinated_prefill: @@ -350,6 +350,7 @@ def _local_prefill_cuda_graph_vote( capture_hidden_mode=None, return_logprob=return_logprob, lora_ineligible=prefill_graph_runner.enable_lora, + is_mixed=mode == ForwardMode.MIXED, batch_max_context_len=( int(local_batch.seq_lens_cpu.max().item()) if prefill_graph_runner.max_context_size is not None diff --git a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py index f980a01a0..f0d3ab989 100644 --- a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py @@ -302,6 +302,15 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): # --- prefill graph config ------------------------------------- prefill_config = get_exec().graph.cuda_graph_config.prefill self.prefill_backend_name = prefill_config.backend + self.prefer_eager_mixed_prefill = ( + self.prefill_backend_name == Backend.BREAKABLE + and get_parallel().enable_dp_attention + and getattr( + model_runner.attn_backend, + "prefer_eager_mixed_prefill_under_dp_attention", + False, + ) + ) # bs in prefill carries the captured shape (token count for # tc_piecewise) — one shape knob per phase. capture_tokens = prefill_config.bs @@ -1199,6 +1208,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): capture_hidden_mode, return_logprob: bool, lora_ineligible: bool = False, + is_mixed: bool = False, batch_max_context_len: Optional[int] = None, ) -> bool: """Rank-local replay eligibility: the single source of truth for @@ -1215,6 +1225,8 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): # schedule-time vote derives this from enable_lora alone. if lora_ineligible: return False + if is_mixed and getattr(self, "prefer_eager_mixed_prefill", False): + return False if input_embeds is not None: return False if replace_embeds is not None: @@ -1295,6 +1307,14 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): forward_batch ) ), + is_mixed=any( + getattr(forward_batch, field, None) == ForwardMode.MIXED + for field in ( + "forward_mode", + "global_forward_mode", + "_original_forward_mode", + ) + ), batch_max_context_len=batch_max_context_len, ): return False diff --git a/test/registered/amd/test_dsv4_hip_bcg_metadata.py b/test/registered/amd/test_dsv4_hip_bcg_metadata.py new file mode 100644 index 000000000..f9d0c48ce --- /dev/null +++ b/test/registered/amd/test_dsv4_hip_bcg_metadata.py @@ -0,0 +1,419 @@ +import dataclasses +import unittest +from types import SimpleNamespace +from unittest import mock + +import torch + +from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import ( + DeepseekV4HipRadixBackend, + DeepseekV4MultiStepBackend, + DSV4AttnMetadata, + DSV4Metadata, + UnifiedKvMetadata, + _match_num_queries, +) +from sglang.srt.utils import is_hip +from sglang.test.ci.ci_register import register_amd_ci + +register_amd_ci(est_time=5, suite="stage-b-test-1-gpu-small-amd-mi35x") + + +@unittest.skipUnless(is_hip(), "DeepSeek V4 HIP radix backend requires ROCm") +class TestDSV4HipBreakableCudaGraphMetadata(unittest.TestCase): + @staticmethod + def _make_core_metadata(base: int) -> DSV4AttnMetadata: + def tensor(offset: int) -> torch.Tensor: + return torch.tensor([base + offset], dtype=torch.int32) + + def fill_optional_tensors(metadata, start: int) -> int: + for metadata_field in dataclasses.fields(metadata): + if "Tensor" not in str(metadata_field.type): + continue + if getattr(metadata, metadata_field.name, None) is None: + setattr(metadata, metadata_field.name, tensor(start)) + start += 1 + return start + + metadata = DSV4AttnMetadata( + page_size=256, + page_table=torch.tensor([[base + 1, base + 2]], dtype=torch.int32), + raw_out_loc=torch.tensor([base + 3], dtype=torch.int32), + cuda_int32_kwargs={"dtype": torch.int32}, + seq_lens_casual=torch.tensor([base + 4], dtype=torch.int32), + positions_casual=torch.tensor([base + 5], dtype=torch.int32), + swa_page_indices=torch.tensor([[base + 6, base + 7]], dtype=torch.int32), + swa_topk_lengths=torch.tensor([base + 8], dtype=torch.int32), + c4_sparse_topk=512, + swa_out_cache_loc=torch.tensor([base + 9], dtype=torch.int32), + unified=UnifiedKvMetadata(), + ) + next_offset = fill_optional_tensors(metadata, 10) + fill_optional_tensors(metadata.unified, next_offset) + metadata.c0_flashmla_metadata = None + metadata.c4_flashmla_metadata = None + metadata.c128_flashmla_metadata = None + return metadata + + def test_backend_opts_into_captured_bcg_metadata(self): + self.assertTrue( + DeepseekV4HipRadixBackend.use_captured_forward_metadata_for_breakable_cuda_graph + ) + self.assertTrue( + DeepseekV4HipRadixBackend.prefer_eager_mixed_prefill_under_dp_attention + ) + + def test_non_unified_metadata_matches_underfilled_bucket(self): + captured = torch.tensor([[1, 2], [3, 4], [5, 6], [7, 8]]) + replay = _match_num_queries(captured, 3, value=-1) + self.assertEqual(replay.tolist(), [[1, 2], [3, 4], [5, 6]]) + + short = torch.tensor([9, 10]) + replay = _match_num_queries(short, 3, value=1) + self.assertEqual(replay.tolist(), [9, 10, 1]) + self.assertIsNone(_match_num_queries(None, 3, value=0)) + + def test_unified_prefill_metadata_pads_to_capture_bucket(self): + backend = object.__new__(DeepseekV4HipRadixBackend) + backend.token_to_kv_pool = SimpleNamespace(unified_swa_window=128) + core = self._make_core_metadata(0) + core.positions_casual = torch.tensor([0, 1, 2, 0], dtype=torch.int32) + core.unified = None + + with ( + mock.patch( + "sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate." + "is_unified_kv_triton", + return_value=True, + ), + mock.patch( + "sglang.srt.layers.attention.deepseek_v4_backend_hip_radix." + "torch.repeat_interleave", + wraps=torch.repeat_interleave, + ) as repeat_interleave, + ): + backend._attach_unified_kv_prefill_meta( + core, + req_pool_indices=torch.tensor([7, 9], dtype=torch.int32), + req_pool_indices_repeated=torch.tensor([7, 9, 9, 9], dtype=torch.int32), + seq_lens=torch.tensor([1, 3], dtype=torch.int32), + extend_seq_lens=torch.tensor([1, 2], dtype=torch.int32), + num_tokens=3, + exact_num_tokens=True, + ) + + repeat_interleave.assert_called_once() + self.assertEqual(repeat_interleave.call_args.kwargs["output_size"], 3) + self.assertEqual(core.unified.pf_state_slot.tolist(), [7, 9, 9, 9]) + self.assertEqual(core.unified.pf_chunk_start.tolist(), [0, 1, 1, 0]) + self.assertEqual(core.unified.pf_cu_q.tolist(), [0, 1, 1, 0]) + self.assertEqual(core.unified.pf_final_pos.tolist(), [0, 2, 2, 128]) + + def test_eager_prefill_marks_host_proven_token_count_exact(self): + backend = object.__new__(DeepseekV4HipRadixBackend) + backend.req_to_token = torch.zeros((2, 8), dtype=torch.int32) + backend.token_to_kv_pool = object() + core = self._make_core_metadata(0) + extend_start_loc = torch.tensor([0, 1], dtype=torch.int32) + backend.make_core_attn_metadata = mock.Mock(return_value=core) + backend._attach_unified_kv_prefill_meta = mock.Mock() + backend.init_forward_metadata_indexer = mock.Mock(return_value=None) + + with ( + mock.patch( + "sglang.kernels.ops.attention.dsv4_attn_metadata_kernels." + "ExpandPrefillCausally.execute", + return_value=SimpleNamespace( + seq_lens_casual=core.seq_lens_casual, + req_pool_indices_repeated=torch.tensor( + [7, 9, 9], dtype=torch.int32 + ), + ), + ) as expand_prefill, + mock.patch( + "sglang.srt.layers.attention.deepseek_v4_backend_hip_radix." + "create_paged_compressor_data", + return_value=None, + ), + ): + backend.init_forward_metadata_prefill( + max_seq_len=4096, + req_pool_indices=torch.tensor([7, 9], dtype=torch.int32), + seq_lens=torch.tensor([1, 3], dtype=torch.int32), + seq_lens_cpu=[1, 3], + out_cache_loc=torch.zeros(3, dtype=torch.int64), + num_tokens=3, + extend_seq_lens=torch.tensor([1, 2], dtype=torch.int32), + extend_seq_lens_cpu=[1, 2], + extend_start_loc=extend_start_loc, + use_prefill_cuda_graph=False, + exact_num_tokens=False, + ) + + self.assertIs( + expand_prefill.call_args.kwargs["extend_start_loc"], extend_start_loc + ) + self.assertTrue( + backend._attach_unified_kv_prefill_meta.call_args.kwargs["exact_num_tokens"] + ) + + def test_prefill_bcg_uses_bucket_sized_gpu_compressor_plans(self): + backend = object.__new__(DeepseekV4HipRadixBackend) + backend.req_to_token = torch.zeros((2, 8), dtype=torch.int32) + backend.token_to_kv_pool = object() + core = self._make_core_metadata(0) + core.positions_casual = torch.tensor([0, 1, 2, 0], dtype=torch.int32) + backend.make_core_attn_metadata = mock.Mock(return_value=core) + backend._attach_unified_kv_prefill_meta = mock.Mock() + backend.init_forward_metadata_indexer = mock.Mock(return_value=None) + + with mock.patch( + "sglang.srt.layers.attention.deepseek_v4_backend_hip_radix." + "create_paged_compressor_data", + side_effect=lambda compress_ratio, **kwargs: (compress_ratio, kwargs), + ) as create_plan: + backend.init_forward_metadata_prefill( + max_seq_len=4096, + req_pool_indices=torch.tensor([7, 9], dtype=torch.int32), + seq_lens=torch.tensor([1, 3], dtype=torch.int32), + seq_lens_cpu=[1, 3], + out_cache_loc=torch.zeros(4, dtype=torch.int64), + num_tokens=3, + extend_seq_lens=torch.tensor([1, 2], dtype=torch.int32), + extend_seq_lens_cpu=[1, 2], + use_prefill_cuda_graph=True, + ) + + self.assertEqual(create_plan.call_count, 2) + for call in create_plan.call_args_list: + self.assertIsNone(call.kwargs["seq_lens_cpu"]) + self.assertIsNone(call.kwargs["extend_lens_cpu"]) + self.assertEqual(call.kwargs["num_q_tokens"], 4) + self.assertTrue(call.kwargs["use_prefill_cuda_graph"]) + + def test_gpu_compressor_plan_invalidates_bucket_tail(self): + from sglang.kernels.ops.attention.dsv4 import CompressorPrefillPlan + from sglang.test.kernels.deepseek_v4.common import make_paged_context + + seq_lens = torch.tensor([1, 3], dtype=torch.int64, device="cuda") + extend_lens = torch.tensor([1, 2], dtype=torch.int64, device="cuda") + for compress_ratio in (4, 128): + with self.subTest(compress_ratio=compress_ratio): + context = make_paged_context(bs=2, compress_ratio=compress_ratio) + plan = CompressorPrefillPlan.generate( + compress_ratio=compress_ratio, + req_pool_indices=context.req_pool_indices, + seq_lens=seq_lens, + extend_lens=extend_lens, + req_to_token=context.req_to_token, + full_to_state=context.full_to_swa, + swa_page_size=context.swa_page_size, + ring_size=context.ring_size, + num_q_tokens=4, + use_cuda_graph=True, + ) + + self.assertEqual(plan.plan_c.shape, (4, 16)) + self.assertEqual(plan.plan_w.shape, (4, 8)) + ragged_ids = plan.plan_w.view(torch.uint32).view(-1, 2)[:, 0] + self.assertEqual(ragged_ids[:3].cpu().tolist(), [0, 1, 2]) + self.assertEqual(int(ragged_ids[3].item()), 0xFFFFFFFF) + + def test_capture_builds_graph_compatible_metadata_and_workspace(self): + capture_metadata = DSV4Metadata(object(), indexer_metadata=None) + backend = object.__new__(DeepseekV4HipRadixBackend) + backend.MAX_SEQ_LEN_FOR_CAPTURE = 4096 + backend._build_forward_metadata = mock.Mock(return_value=capture_metadata) + backend.init_forward_metadata_in_graph = mock.Mock() + backend._refresh_fp4_prefill_workspace = mock.Mock() + forward_batch = SimpleNamespace(name="capture") + + result = backend.init_forward_metadata_for_breakable_cuda_graph_capture( + forward_batch + ) + + backend._build_forward_metadata.assert_called_once_with( + forward_batch, + max_seq_len_override=backend.MAX_SEQ_LEN_FOR_CAPTURE, + use_prefill_cuda_graph=True, + ) + backend.init_forward_metadata_in_graph.assert_called_once_with(forward_batch) + backend._refresh_fp4_prefill_workspace.assert_called_once_with(forward_batch) + self.assertIs(result, capture_metadata) + self.assertIs(backend.forward_metadata, capture_metadata) + + def test_refresh_preserves_captured_hip_tensor_storage(self): + capture_workspace = object() + capture_metadata = DSV4Metadata( + self._make_core_metadata(0), + indexer_metadata=None, + fp4_prefill_workspace=capture_workspace, + fp4_k_write_metadata=( + torch.tensor([14], dtype=torch.int64), + torch.tensor([15], dtype=torch.int64), + ), + fp4_q_positions=torch.tensor([16], dtype=torch.int64), + ) + replay_metadata = DSV4Metadata( + self._make_core_metadata(100), + indexer_metadata=None, + fp4_k_write_metadata=( + torch.tensor([114], dtype=torch.int64), + torch.tensor([115], dtype=torch.int64), + ), + fp4_q_positions=torch.tensor([116], dtype=torch.int64), + ) + capture_core = capture_metadata.core_attn_metadata + replay_core = replay_metadata.core_attn_metadata + captured_core_tensors = { + field.name: getattr(capture_core, field.name) + for field in dataclasses.fields(capture_core) + if torch.is_tensor(getattr(capture_core, field.name)) + } + captured_unified_tensors = { + field.name: getattr(capture_core.unified, field.name) + for field in dataclasses.fields(capture_core.unified) + if torch.is_tensor(getattr(capture_core.unified, field.name)) + } + captured_fp4_tensors = { + "fp4_k_positions": capture_metadata.fp4_k_write_metadata[0], + "fp4_k_slots": capture_metadata.fp4_k_write_metadata[1], + "fp4_q_positions": capture_metadata.fp4_q_positions, + } + expected_core_tensors = { + name: getattr(replay_core, name).clone() for name in captured_core_tensors + } + expected_unified_tensors = { + name: getattr(replay_core.unified, name).clone() + for name in captured_unified_tensors + } + replay_fp4_tensors = { + "fp4_k_positions": replay_metadata.fp4_k_write_metadata[0], + "fp4_k_slots": replay_metadata.fp4_k_write_metadata[1], + "fp4_q_positions": replay_metadata.fp4_q_positions, + } + expected_fp4_tensors = { + name: tensor.clone() for name, tensor in replay_fp4_tensors.items() + } + + capture_metadata.refresh_for_breakable_cuda_graph_replay_(replay_metadata) + + for field_name, captured_tensor in captured_core_tensors.items(): + current = getattr(capture_core, field_name) + self.assertIs(current, captured_tensor, field_name) + self.assertTrue( + torch.equal(current, expected_core_tensors[field_name]), field_name + ) + self.assertTrue( + torch.equal( + getattr(replay_core, field_name), + expected_core_tensors[field_name], + ), + f"{field_name} replay source", + ) + for field_name, captured_tensor in captured_unified_tensors.items(): + current = getattr(capture_core.unified, field_name) + self.assertIs(current, captured_tensor, field_name) + self.assertTrue( + torch.equal(current, expected_unified_tensors[field_name]), + field_name, + ) + self.assertTrue( + torch.equal( + getattr(replay_core.unified, field_name), + expected_unified_tensors[field_name], + ), + f"{field_name} replay source", + ) + + current_fp4_tensors = { + "fp4_k_positions": capture_metadata.fp4_k_write_metadata[0], + "fp4_k_slots": capture_metadata.fp4_k_write_metadata[1], + "fp4_q_positions": capture_metadata.fp4_q_positions, + } + for name, captured_tensor in captured_fp4_tensors.items(): + self.assertIs(current_fp4_tensors[name], captured_tensor) + self.assertTrue( + torch.equal(captured_tensor, expected_fp4_tensors[name]), name + ) + self.assertTrue( + torch.equal(replay_fp4_tensors[name], expected_fp4_tensors[name]), + f"{name} replay source", + ) + self.assertIs(capture_metadata.fp4_prefill_workspace, capture_workspace) + + def test_replay_refreshes_captured_metadata_and_workspace(self): + capture_metadata = DSV4Metadata(object(), indexer_metadata=None) + replay_metadata = DSV4Metadata(object(), indexer_metadata=None) + capture_metadata.refresh_for_breakable_cuda_graph_replay_ = mock.Mock() + + backend = object.__new__(DeepseekV4HipRadixBackend) + backend.MAX_SEQ_LEN_FOR_CAPTURE = 4096 + backend._build_forward_metadata = mock.Mock(return_value=replay_metadata) + backend.init_forward_metadata_in_graph = mock.Mock() + backend._refresh_fp4_prefill_workspace = mock.Mock() + + forward_batch = SimpleNamespace(name="live") + static_forward_batch = SimpleNamespace(name="static") + backend.prepare_forward_metadata_for_breakable_cuda_graph_replay( + capture_metadata, + forward_batch, + static_forward_batch=static_forward_batch, + ) + + backend._build_forward_metadata.assert_called_once_with( + static_forward_batch, + max_seq_len_override=backend.MAX_SEQ_LEN_FOR_CAPTURE, + use_prefill_cuda_graph=True, + ) + backend.init_forward_metadata_in_graph.assert_called_once_with( + static_forward_batch + ) + capture_metadata.refresh_for_breakable_cuda_graph_replay_.assert_called_once_with( + replay_metadata + ) + backend._refresh_fp4_prefill_workspace.assert_called_once_with( + static_forward_batch + ) + self.assertIs(backend.forward_metadata, capture_metadata) + + def test_multistep_backend_forwards_bcg_metadata_hooks(self): + backend = object.__new__(DeepseekV4MultiStepBackend) + backend.speculative_num_steps = 3 + backend.attn_backends = [mock.Mock(), mock.Mock(), mock.Mock()] + forward_batch = SimpleNamespace(name="live") + static_forward_batch = SimpleNamespace(name="static") + capture_metadata = [object(), object()] + + for index, child in enumerate(backend.attn_backends[:-1]): + child.init_forward_metadata_for_breakable_cuda_graph_capture.return_value = f"capture-{index}" + + captured = backend.init_forward_metadata_for_breakable_cuda_graph_capture( + forward_batch + ) + self.assertEqual(captured, ["capture-0", "capture-1"]) + + backend.prepare_forward_metadata_for_breakable_cuda_graph_replay( + capture_metadata, + forward_batch, + static_forward_batch=static_forward_batch, + ) + for index, child in enumerate(backend.attn_backends[:-1]): + child.init_forward_metadata_for_breakable_cuda_graph_capture.assert_called_once_with( + forward_batch + ) + child.prepare_forward_metadata_for_breakable_cuda_graph_replay.assert_called_once_with( + capture_metadata[index], + forward_batch, + static_forward_batch=static_forward_batch, + ) + backend.attn_backends[ + -1 + ].init_forward_metadata_for_breakable_cuda_graph_capture.assert_not_called() + backend.attn_backends[ + -1 + ].prepare_forward_metadata_for_breakable_cuda_graph_replay.assert_not_called() + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/managers/scheduler_components/test_dp_attn.py b/test/registered/unit/managers/scheduler_components/test_dp_attn.py index 2dbd1cfdd..5696248f6 100644 --- a/test/registered/unit/managers/scheduler_components/test_dp_attn.py +++ b/test/registered/unit/managers/scheduler_components/test_dp_attn.py @@ -131,5 +131,45 @@ class TestDecodeToExtendConversionVote(CustomTestCase): self.assertFalse(self._vote(beam=True)) +class TestPrefillCudaGraphVote(CustomTestCase): + def _vote(self, mode): + runner = Mock(spec=dp_attn.PrefillCudaGraphRunner) + runner.enable_lora = False + runner.max_context_size = None + runner.can_replay_locally.return_value = True + batch = SimpleNamespace( + forward_mode=mode, + extend_num_tokens=4, + input_embeds=None, + replace_embeds=None, + prefix_lens=[1, 1], + return_logprob=False, + batch_size=lambda: 2, + ) + vote = dp_attn._local_prefill_cuda_graph_vote( + local_batch=batch, + prefill_graph_runner=runner, + coordinated_prefill=True, + breakable_prefill=True, + spec_algorithm=SpeculativeAlgorithm.NONE, + model_config=object(), + ) + return vote, runner + + def test_extend_batch_votes_for_prefill_graph(self): + vote, runner = self._vote(ForwardMode.EXTEND) + + self.assertTrue(vote) + runner.can_replay_locally.assert_called_once() + self.assertFalse(runner.can_replay_locally.call_args.kwargs["is_mixed"]) + + def test_mixed_batch_delegates_to_runner_policy(self): + vote, runner = self._vote(ForwardMode.MIXED) + + self.assertTrue(vote) + runner.can_replay_locally.assert_called_once() + self.assertTrue(runner.can_replay_locally.call_args.kwargs["is_mixed"]) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/model_executor/runner/test_prefill_cuda_graph_padding.py b/test/registered/unit/model_executor/runner/test_prefill_cuda_graph_padding.py index 3ff162b08..f1474f3dc 100644 --- a/test/registered/unit/model_executor/runner/test_prefill_cuda_graph_padding.py +++ b/test/registered/unit/model_executor/runner/test_prefill_cuda_graph_padding.py @@ -31,19 +31,22 @@ class TestPrefillCudaGraphPadding(CustomTestCase): runner._capture_chunked_prefix = False runner.prefill_backend_name = Backend.TC_PIECEWISE runner.has_mha_companion_layers = False + runner.prefer_eager_mixed_prefill = False runner.capture_hidden_mode = CaptureHiddenMode.NULL runner.capture_num_tokens = [4, 16] runner.max_context_size = None runner.max_num_tokens = 16 return runner - def _make_forward_batch(self, num_tokens): + def _make_forward_batch(self, num_tokens, mode=ForwardMode.EXTEND): return SimpleNamespace( batch_size=1, input_embeds=None, replace_embeds=None, mm_inputs=None, - forward_mode=ForwardMode.EXTEND, + forward_mode=mode, + global_forward_mode=None, + _original_forward_mode=None, capture_hidden_mode=CaptureHiddenMode.NULL, global_num_tokens_cpu=None, return_logprob=False, @@ -63,6 +66,14 @@ class TestPrefillCudaGraphPadding(CustomTestCase): self.assertTrue(runner.can_run_graph(self._make_forward_batch(8))) + def test_mixed_batch_uses_scoped_runner_policy(self): + runner = self._make_runner() + batch = self._make_forward_batch(8, mode=ForwardMode.MIXED) + + self.assertTrue(runner.can_run_graph(batch)) + runner.prefer_eager_mixed_prefill = True + self.assertFalse(runner.can_run_graph(batch)) + def test_replay_snapshot_uses_padded_token_count(self): runner = self._make_runner() runner.use_captured_attn_metadata = False