From 21289cfd50a8cf7550c1786ac71733bec9396d4d Mon Sep 17 00:00:00 2001 From: Xinyi Song Date: Sat, 12 Sep 2026 15:21:43 -0700 Subject: [PATCH] [AMD] Fix Dspark accept length and reduce host bubble on DSV4 (#39116) --- .../deepseek_v4_backend_hip_radix.py | 178 +++++++++--------- .../amd/test_deepseek_v4_pro_fp4_dspark.py | 61 ++++++ 2 files changed, 155 insertions(+), 84 deletions(-) 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 9ac37dd0a..8e04117c8 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 @@ -10,7 +10,6 @@ from typing import ( List, Literal, Optional, - Tuple, TypeVar, Union, ) @@ -141,10 +140,11 @@ class UnifiedKvMetadata: "verify_store_state_slot", "c4_out_loc", "c128_out_loc", + # Captured store_cache reads swa_loc by address, and the eager + # target-verify path builds it outside the graph. + "swa_loc", ], - # swa_loc is recomputed each forward (recorded inside cuda graphs), - # so it is rebound rather than copied across replays. - assign_fields=["swa_loc"], + assign_fields=[], ) @@ -482,9 +482,8 @@ class DeepseekV4HipRadixBackend( self.speculative_num_steps = speculative_num_steps self.speculative_num_draft_tokens: int = get_spec().speculative_num_draft_tokens self.is_draft_worker = getattr(model_runner, "is_draft_worker", False) - self.is_dspark_draft = ( - self.is_draft_worker and model_runner.spec_algorithm.is_dspark() - ) + self.is_dspark = model_runner.spec_algorithm.is_dspark() + self.is_dspark_draft = self.is_draft_worker and self.is_dspark self.target_verify_num_draft_tokens = self.speculative_num_draft_tokens if self.is_dspark_draft: assert self.speculative_num_draft_tokens is not None @@ -493,6 +492,15 @@ class DeepseekV4HipRadixBackend( # CUDA-side convention gamma + 1, so use an explicit effective value # instead of mutating speculative_num_draft_tokens in place. self.target_verify_num_draft_tokens = self.speculative_num_draft_tokens - 1 + # Past MAX_FUSED_ROWS the fp4 schedule falls back to AITER's preamble, + # which frees the scratch its kernels read -- not capture-safe. + self._fp4_graph_row_limit: Optional[int] = None + if self.enable_deepseek_v4_fp4_indexer and self.speculative_num_steps == 0: + from sglang.kernels.ops.attention.dsv4.fp4_indexer_schedule_hip import ( + MAX_FUSED_ROWS, + ) + + self._fp4_graph_row_limit = MAX_FUSED_ROWS self.speculative_step_id = speculative_step_id self.forward_metadata: Union[ DSV4Metadata, @@ -545,32 +553,29 @@ class DeepseekV4HipRadixBackend( compress_gpu_plan: bool = False, extend_start_loc: Optional[torch.Tensor] = None, attach_decode_streams: bool = False, + # Whether num_tokens == sum(extend_seq_lens) exactly, which lets the + # token map skip an implicit D2H. + exact_num_tokens: bool = True, ) -> DSV4Metadata: - if extend_start_loc is not None: - from sglang.kernels.ops.attention.dsv4_attn_metadata_kernels import ( - ExpandPrefillCausally, - ) + from sglang.kernels.ops.attention.dsv4_attn_metadata_kernels import ( + ExpandPrefillCausally, + ) - _expanded = ExpandPrefillCausally.execute( - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - extend_seq_lens=extend_seq_lens, - extend_start_loc=extend_start_loc, - seq_lens_cpu=None, - extend_seq_lens_cpu=None, - num_tokens=num_tokens, - padded_num_tokens=out_cache_loc.shape[0], - ) - seq_lens_casual = _expanded.seq_lens_casual - req_pool_indices_repeated = _expanded.req_pool_indices_repeated - else: - seq_lens_casual, req_pool_indices_repeated = self.expand_prefill_casually( - num_tokens=num_tokens, - seq_lens=seq_lens_cpu, - extend_seq_lens=extend_seq_lens_cpu, - req_pool_indices=req_pool_indices, - padded_num_tokens=out_cache_loc.shape[0], - ) + # extend_start_loc and the CPU mirrors below only feed the torch + # fallback; the triton kernel cumsums extend_seq_lens on device, so + # every caller can share it instead of dropping to a per-request loop. + _expanded = ExpandPrefillCausally.execute( + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + extend_seq_lens=extend_seq_lens, + extend_start_loc=extend_start_loc, + seq_lens_cpu=seq_lens_cpu, + extend_seq_lens_cpu=extend_seq_lens_cpu, + num_tokens=num_tokens, + padded_num_tokens=out_cache_loc.shape[0], + ) + seq_lens_casual = _expanded.seq_lens_casual + req_pool_indices_repeated = _expanded.req_pool_indices_repeated core_attn_metadata = self.make_core_attn_metadata( req_to_token=self.req_to_token, req_pool_indices_repeated=req_pool_indices_repeated, @@ -586,7 +591,7 @@ class DeepseekV4HipRadixBackend( seq_lens, extend_seq_lens, num_tokens, - need_compress=need_compress, + exact_num_tokens=exact_num_tokens, ) if attach_decode_streams: # Target-verify runs through the unified_kv DECODE kernel, so build @@ -648,8 +653,29 @@ class DeepseekV4HipRadixBackend( seq_lens_cpu: Optional[List[int]] = None, ragged_layout=None, ) -> Union[DSV4Metadata, DSV4RawVerifyMetadata]: - # HIP path: build target-verify metadata eagerly. The raw/lazy-upgrade route can - # hit planner invariants during graph capture for DSV4+EAGLE. + # DSPARK verifies a uniform num_draft block, exactly what + # make_forward_metadata_from_raw_verify expands, so the build can be + # deferred into the graph. Graph path only: the upgrade sizes its page + # table by MAX_SEQ_LEN_FOR_CAPTURE, far wider than the live max_seq_len + # an eager caller passes. EAGLE and ragged layouts stay eager -- no raw + # expansion, and EAGLE's fixed-tier plan trips planner invariants. + if ( + use_prefill_cuda_graph + and self.is_dspark + and ragged_layout is None + and out_cache_loc is not None + # Oversized batches keep the eager build; see _fp4_graph_row_limit. + and ( + self._fp4_graph_row_limit is None + or self.target_verify_num_draft_tokens * len(seq_lens) + <= self._fp4_graph_row_limit + ) + ): + return DSV4RawVerifyMetadata( + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + out_cache_loc=out_cache_loc, + ) if seq_lens_cpu is None: seq_lens_cpu = seq_lens.tolist() return self.init_forward_metadata_target_verify_old( @@ -692,8 +718,11 @@ class DeepseekV4HipRadixBackend( num_tokens = ragged_layout.total_verify_tokens if num_tokens is None: num_tokens = int(verify_lens_dev.sum().item()) + exact_num_tokens = True else: num_tokens = int(num_tokens) + # Padded tier: num_tokens >= sum(verify_lens) + exact_num_tokens = False extend_seq_lens_cpu = None seq_lens_cpu = None else: @@ -704,6 +733,7 @@ class DeepseekV4HipRadixBackend( extend_seq_lens_cpu = [self.target_verify_num_draft_tokens] * batch_size num_tokens = self.target_verify_num_draft_tokens * batch_size extend_seq_lens = self._move_to_device(extend_seq_lens_cpu) + exact_num_tokens = True if out_cache_loc is None: out_cache_loc = seq_lens.new_zeros(num_tokens) return self.init_forward_metadata_prefill( @@ -720,6 +750,7 @@ class DeepseekV4HipRadixBackend( compress_gpu_plan=ragged_layout is not None, extend_start_loc=extend_start_loc, attach_decode_streams=True, + exact_num_tokens=exact_num_tokens, ) def make_forward_metadata_from_raw_verify( @@ -753,6 +784,18 @@ class DeepseekV4HipRadixBackend( out_loc=out_cache_loc, need_compress=True, ) + # extend_seq_lens is uniform here (seq_lens already carries the draft + # block, so the minimum above cannot trim it), hence an exact token count. + self._attach_unified_kv_prefill_meta( + core_attn_metadata, + req_pool_indices, + seq_lens, + extend_seq_lens, + num_draft_tokens * bs, + ) + self._attach_unified_kv_decode_streams( + core_attn_metadata, req_pool_indices_repeated + ) indexer_metadata = self.init_forward_metadata_indexer(core_attn_metadata) create = functools.partial( create_paged_compressor_data, @@ -843,7 +886,8 @@ class DeepseekV4HipRadixBackend( def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: # Raw metadata must be materialized inside the graph to refresh on replay. - if isinstance(self.forward_metadata, DSV4RawVerifyMetadata): + upgraded_verify = isinstance(self.forward_metadata, DSV4RawVerifyMetadata) + if upgraded_verify: self.forward_metadata = self.make_forward_metadata_from_raw_verify( raw_metadata=self.forward_metadata, ) @@ -894,6 +938,12 @@ class DeepseekV4HipRadixBackend( torch.int64 ) + if upgraded_verify: + # The out-graph refresh saw raw metadata and skipped. Without this + # the logits kernel builds its own schedule -- the variant that + # frees the scratch it reads, which every replay would re-read. + self._refresh_fp4_prefill_workspace(forward_batch) + # Decode's schedule builder is capture-safe because the workspace pins # the scratch it reads, so it can stay next to the metadata it consumes. # Prefill/target-verify cannot; see _refresh_fp4_prefill_workspace. @@ -1159,6 +1209,7 @@ class DeepseekV4HipRadixBackend( extend_seq_lens=extend_seq_lens, extend_seq_lens_cpu=extend_seq_lens_cpu, need_compress=not is_draft, + exact_num_tokens=is_draft, ) else: raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}") @@ -1276,7 +1327,7 @@ class DeepseekV4HipRadixBackend( seq_lens: torch.Tensor, extend_seq_lens: torch.Tensor, num_tokens: int, - need_compress: bool = True, + exact_num_tokens: bool = True, ) -> None: from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( is_unified_kv_triton, @@ -1289,19 +1340,13 @@ class DeepseekV4HipRadixBackend( seq_lens = seq_lens.to(torch.int64) extend_seq_lens = extend_seq_lens.to(torch.int64) # token -> req index (length L = sum(extend_seq_lens)). - # output_size skips the implicit sum() D2H on draft-extend. dropping it on the - # target-extend path triggers a GPU memory access fault. - if need_compress: - bid = torch.repeat_interleave( - torch.arange(bs, device=device, dtype=torch.int64), - extend_seq_lens, - ) - else: - bid = torch.repeat_interleave( - torch.arange(bs, device=device, dtype=torch.int64), - extend_seq_lens, - output_size=num_tokens, - ) + # output_size skips the implicit sum() D2H, but it must equal L: + # exact_num_tokens tells whether num_tokens does. + bid = torch.repeat_interleave( + torch.arange(bs, device=device, dtype=torch.int64), + extend_seq_lens, + output_size=num_tokens if exact_num_tokens else None, + ) if core.unified is None: core.unified = UnifiedKvMetadata() core.unified.pf_state_slot = req_pool_indices[bid] @@ -1680,41 +1725,6 @@ class DeepseekV4HipRadixBackend( raise NotImplementedError("ragged attention") - def expand_prefill_casually( - self, - num_tokens: int, - seq_lens: List[int], - extend_seq_lens: List[int], - req_pool_indices: torch.Tensor, - padded_num_tokens: Optional[int], - ) -> Tuple[torch.Tensor, torch.Tensor]: - seq_lens_casual = torch.empty(num_tokens, **self.cuda_int32_kwargs) - idx_to_req_repeated = torch.empty(num_tokens, **self.cuda_int32_kwargs) - offset = 0 - for i, (kv_len, qo_len) in enumerate(zip(seq_lens, extend_seq_lens)): - out = seq_lens_casual[offset : offset + qo_len] - offset += qo_len - torch.arange(kv_len - qo_len + 1, kv_len + 1, out=out) - idx_to_req_repeated[offset - qo_len : offset].fill_(i) - - assert offset == num_tokens - req_pool_indices_repeated = req_pool_indices[idx_to_req_repeated] - - if padded_num_tokens is not None and padded_num_tokens > num_tokens: - pad_size = padded_num_tokens - num_tokens - 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(), - ) - - return seq_lens_casual, req_pool_indices_repeated - def expand_extend_with_same_length( self, bs: int, diff --git a/test/registered/amd/test_deepseek_v4_pro_fp4_dspark.py b/test/registered/amd/test_deepseek_v4_pro_fp4_dspark.py index d3a9d4fee..54d679b15 100644 --- a/test/registered/amd/test_deepseek_v4_pro_fp4_dspark.py +++ b/test/registered/amd/test_deepseek_v4_pro_fp4_dspark.py @@ -11,12 +11,19 @@ Registry: nightly-amd-8-gpu-mi35x-deepseek-v4-pro-dspark suite import os import unittest from types import SimpleNamespace +from unittest import mock import requests import torch +from sglang.kernels.ops.attention.dsv4.fp4_indexer_schedule_hip import MAX_FUSED_ROWS from sglang.kernels.ops.attention.dsv4.unified_kv_kernels import runtime from sglang.kernels.ops.speculative.dspark import dspark_verify_window +from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import ( + DeepseekV4HipRadixBackend, + DSV4RawVerifyMetadata, + UnifiedKvMetadata, +) from sglang.srt.utils import kill_process_tree from sglang.test.ci.ci_register import register_amd_ci from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k @@ -60,6 +67,60 @@ FP4_ENV_VARS = { class TestDSparkUnifiedKVKernelsAMD(CustomTestCase): + def test_unified_metadata_copy_updates_captured_swa_loc_in_place(self): + captured_swa_loc = torch.tensor([3, 5, 7], device=DEVICE, dtype=torch.int32) + replay_swa_loc = torch.tensor([11, 13, 17], device=DEVICE, dtype=torch.int32) + captured = UnifiedKvMetadata(swa_loc=captured_swa_loc) + replay = UnifiedKvMetadata(swa_loc=replay_swa_loc) + + captured.copy_(replay) + + self.assertIs(captured.swa_loc, captured_swa_loc) + self.assertTrue(torch.equal(captured.swa_loc, replay_swa_loc)) + + def test_dspark_verify_metadata_graph_routing(self): + num_draft_tokens = 7 + max_fused_bs = MAX_FUSED_ROWS // num_draft_tokens + cases = ( + ("eligible", True, None, max_fused_bs, True), + ("fp4_row_limit", True, None, max_fused_bs + 1, False), + ("eagle", False, None, max_fused_bs, False), + ("ragged", True, object(), max_fused_bs, False), + ) + + for name, is_dspark, ragged_layout, bs, expect_raw in cases: + with self.subTest(name=name): + backend = DeepseekV4HipRadixBackend.__new__(DeepseekV4HipRadixBackend) + backend.is_dspark = is_dspark + backend._fp4_graph_row_limit = MAX_FUSED_ROWS + backend.target_verify_num_draft_tokens = num_draft_tokens + eager_metadata = object() + backend.init_forward_metadata_target_verify_old = mock.Mock( + return_value=eager_metadata + ) + + req_pool_indices = torch.arange(bs, device=DEVICE, dtype=torch.int32) + seq_lens = torch.ones(bs, device=DEVICE, dtype=torch.int32) + out_cache_loc = torch.zeros( + bs * num_draft_tokens, device=DEVICE, dtype=torch.int32 + ) + result = backend.init_forward_metadata_target_verify( + max_seq_len=128, + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + out_cache_loc=out_cache_loc, + use_prefill_cuda_graph=True, + seq_lens_cpu=[1] * bs, + ragged_layout=ragged_layout, + ) + + if expect_raw: + self.assertIsInstance(result, DSV4RawVerifyMetadata) + backend.init_forward_metadata_target_verify_old.assert_not_called() + else: + self.assertIs(result, eager_metadata) + backend.init_forward_metadata_target_verify_old.assert_called_once() + def test_build_unified_commit_inject_layout(self): stride, ring_stride = 7, 128 req_pool_indices = torch.tensor([3, 0, 5, 1], device=DEVICE, dtype=torch.int32)