[AMD] Fix Dspark accept length and reduce host bubble on DSV4 (#39116)

This commit is contained in:
Xinyi Song
2026-09-12 15:21:43 -07:00
committed by GitHub
parent b5a2aebc7e
commit 21289cfd50
2 changed files with 155 additions and 84 deletions
@@ -10,7 +10,6 @@ from typing import (
List, List,
Literal, Literal,
Optional, Optional,
Tuple,
TypeVar, TypeVar,
Union, Union,
) )
@@ -141,10 +140,11 @@ class UnifiedKvMetadata:
"verify_store_state_slot", "verify_store_state_slot",
"c4_out_loc", "c4_out_loc",
"c128_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), assign_fields=[],
# so it is rebound rather than copied across replays.
assign_fields=["swa_loc"],
) )
@@ -482,9 +482,8 @@ class DeepseekV4HipRadixBackend(
self.speculative_num_steps = speculative_num_steps self.speculative_num_steps = speculative_num_steps
self.speculative_num_draft_tokens: int = get_spec().speculative_num_draft_tokens 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_draft_worker = getattr(model_runner, "is_draft_worker", False)
self.is_dspark_draft = ( self.is_dspark = model_runner.spec_algorithm.is_dspark()
self.is_draft_worker and 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 self.target_verify_num_draft_tokens = self.speculative_num_draft_tokens
if self.is_dspark_draft: if self.is_dspark_draft:
assert self.speculative_num_draft_tokens is not None 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 # CUDA-side convention gamma + 1, so use an explicit effective value
# instead of mutating speculative_num_draft_tokens in place. # instead of mutating speculative_num_draft_tokens in place.
self.target_verify_num_draft_tokens = self.speculative_num_draft_tokens - 1 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.speculative_step_id = speculative_step_id
self.forward_metadata: Union[ self.forward_metadata: Union[
DSV4Metadata, DSV4Metadata,
@@ -545,32 +553,29 @@ class DeepseekV4HipRadixBackend(
compress_gpu_plan: bool = False, compress_gpu_plan: bool = False,
extend_start_loc: Optional[torch.Tensor] = None, extend_start_loc: Optional[torch.Tensor] = None,
attach_decode_streams: bool = False, 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: ) -> DSV4Metadata:
if extend_start_loc is not None: from sglang.kernels.ops.attention.dsv4_attn_metadata_kernels import (
from sglang.kernels.ops.attention.dsv4_attn_metadata_kernels import ( ExpandPrefillCausally,
ExpandPrefillCausally, )
)
_expanded = ExpandPrefillCausally.execute( # extend_start_loc and the CPU mirrors below only feed the torch
req_pool_indices=req_pool_indices, # fallback; the triton kernel cumsums extend_seq_lens on device, so
seq_lens=seq_lens, # every caller can share it instead of dropping to a per-request loop.
extend_seq_lens=extend_seq_lens, _expanded = ExpandPrefillCausally.execute(
extend_start_loc=extend_start_loc, req_pool_indices=req_pool_indices,
seq_lens_cpu=None, seq_lens=seq_lens,
extend_seq_lens_cpu=None, extend_seq_lens=extend_seq_lens,
num_tokens=num_tokens, extend_start_loc=extend_start_loc,
padded_num_tokens=out_cache_loc.shape[0], seq_lens_cpu=seq_lens_cpu,
) extend_seq_lens_cpu=extend_seq_lens_cpu,
seq_lens_casual = _expanded.seq_lens_casual num_tokens=num_tokens,
req_pool_indices_repeated = _expanded.req_pool_indices_repeated padded_num_tokens=out_cache_loc.shape[0],
else: )
seq_lens_casual, req_pool_indices_repeated = self.expand_prefill_casually( seq_lens_casual = _expanded.seq_lens_casual
num_tokens=num_tokens, req_pool_indices_repeated = _expanded.req_pool_indices_repeated
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],
)
core_attn_metadata = self.make_core_attn_metadata( core_attn_metadata = self.make_core_attn_metadata(
req_to_token=self.req_to_token, req_to_token=self.req_to_token,
req_pool_indices_repeated=req_pool_indices_repeated, req_pool_indices_repeated=req_pool_indices_repeated,
@@ -586,7 +591,7 @@ class DeepseekV4HipRadixBackend(
seq_lens, seq_lens,
extend_seq_lens, extend_seq_lens,
num_tokens, num_tokens,
need_compress=need_compress, exact_num_tokens=exact_num_tokens,
) )
if attach_decode_streams: if attach_decode_streams:
# Target-verify runs through the unified_kv DECODE kernel, so build # Target-verify runs through the unified_kv DECODE kernel, so build
@@ -648,8 +653,29 @@ class DeepseekV4HipRadixBackend(
seq_lens_cpu: Optional[List[int]] = None, seq_lens_cpu: Optional[List[int]] = None,
ragged_layout=None, ragged_layout=None,
) -> Union[DSV4Metadata, DSV4RawVerifyMetadata]: ) -> Union[DSV4Metadata, DSV4RawVerifyMetadata]:
# HIP path: build target-verify metadata eagerly. The raw/lazy-upgrade route can # DSPARK verifies a uniform num_draft block, exactly what
# hit planner invariants during graph capture for DSV4+EAGLE. # 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: if seq_lens_cpu is None:
seq_lens_cpu = seq_lens.tolist() seq_lens_cpu = seq_lens.tolist()
return self.init_forward_metadata_target_verify_old( return self.init_forward_metadata_target_verify_old(
@@ -692,8 +718,11 @@ class DeepseekV4HipRadixBackend(
num_tokens = ragged_layout.total_verify_tokens num_tokens = ragged_layout.total_verify_tokens
if num_tokens is None: if num_tokens is None:
num_tokens = int(verify_lens_dev.sum().item()) num_tokens = int(verify_lens_dev.sum().item())
exact_num_tokens = True
else: else:
num_tokens = int(num_tokens) num_tokens = int(num_tokens)
# Padded tier: num_tokens >= sum(verify_lens)
exact_num_tokens = False
extend_seq_lens_cpu = None extend_seq_lens_cpu = None
seq_lens_cpu = None seq_lens_cpu = None
else: else:
@@ -704,6 +733,7 @@ class DeepseekV4HipRadixBackend(
extend_seq_lens_cpu = [self.target_verify_num_draft_tokens] * batch_size extend_seq_lens_cpu = [self.target_verify_num_draft_tokens] * batch_size
num_tokens = 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) extend_seq_lens = self._move_to_device(extend_seq_lens_cpu)
exact_num_tokens = True
if out_cache_loc is None: if out_cache_loc is None:
out_cache_loc = seq_lens.new_zeros(num_tokens) out_cache_loc = seq_lens.new_zeros(num_tokens)
return self.init_forward_metadata_prefill( return self.init_forward_metadata_prefill(
@@ -720,6 +750,7 @@ class DeepseekV4HipRadixBackend(
compress_gpu_plan=ragged_layout is not None, compress_gpu_plan=ragged_layout is not None,
extend_start_loc=extend_start_loc, extend_start_loc=extend_start_loc,
attach_decode_streams=True, attach_decode_streams=True,
exact_num_tokens=exact_num_tokens,
) )
def make_forward_metadata_from_raw_verify( def make_forward_metadata_from_raw_verify(
@@ -753,6 +784,18 @@ class DeepseekV4HipRadixBackend(
out_loc=out_cache_loc, out_loc=out_cache_loc,
need_compress=True, 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) indexer_metadata = self.init_forward_metadata_indexer(core_attn_metadata)
create = functools.partial( create = functools.partial(
create_paged_compressor_data, create_paged_compressor_data,
@@ -843,7 +886,8 @@ class DeepseekV4HipRadixBackend(
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None:
# Raw metadata must be materialized inside the graph to refresh on replay. # 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( self.forward_metadata = self.make_forward_metadata_from_raw_verify(
raw_metadata=self.forward_metadata, raw_metadata=self.forward_metadata,
) )
@@ -894,6 +938,12 @@ class DeepseekV4HipRadixBackend(
torch.int64 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 # 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. # the scratch it reads, so it can stay next to the metadata it consumes.
# Prefill/target-verify cannot; see _refresh_fp4_prefill_workspace. # Prefill/target-verify cannot; see _refresh_fp4_prefill_workspace.
@@ -1159,6 +1209,7 @@ class DeepseekV4HipRadixBackend(
extend_seq_lens=extend_seq_lens, extend_seq_lens=extend_seq_lens,
extend_seq_lens_cpu=extend_seq_lens_cpu, extend_seq_lens_cpu=extend_seq_lens_cpu,
need_compress=not is_draft, need_compress=not is_draft,
exact_num_tokens=is_draft,
) )
else: else:
raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}") raise NotImplementedError(f"unsupported mode {forward_batch.forward_mode=}")
@@ -1276,7 +1327,7 @@ class DeepseekV4HipRadixBackend(
seq_lens: torch.Tensor, seq_lens: torch.Tensor,
extend_seq_lens: torch.Tensor, extend_seq_lens: torch.Tensor,
num_tokens: int, num_tokens: int,
need_compress: bool = True, exact_num_tokens: bool = True,
) -> None: ) -> None:
from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import (
is_unified_kv_triton, is_unified_kv_triton,
@@ -1289,19 +1340,13 @@ class DeepseekV4HipRadixBackend(
seq_lens = seq_lens.to(torch.int64) seq_lens = seq_lens.to(torch.int64)
extend_seq_lens = extend_seq_lens.to(torch.int64) extend_seq_lens = extend_seq_lens.to(torch.int64)
# token -> req index (length L = sum(extend_seq_lens)). # token -> req index (length L = sum(extend_seq_lens)).
# output_size skips the implicit sum() D2H on draft-extend. dropping it on the # output_size skips the implicit sum() D2H, but it must equal L:
# target-extend path triggers a GPU memory access fault. # exact_num_tokens tells whether num_tokens does.
if need_compress: bid = torch.repeat_interleave(
bid = torch.repeat_interleave( torch.arange(bs, device=device, dtype=torch.int64),
torch.arange(bs, device=device, dtype=torch.int64), extend_seq_lens,
extend_seq_lens, output_size=num_tokens if exact_num_tokens else None,
) )
else:
bid = torch.repeat_interleave(
torch.arange(bs, device=device, dtype=torch.int64),
extend_seq_lens,
output_size=num_tokens,
)
if core.unified is None: if core.unified is None:
core.unified = UnifiedKvMetadata() core.unified = UnifiedKvMetadata()
core.unified.pf_state_slot = req_pool_indices[bid] core.unified.pf_state_slot = req_pool_indices[bid]
@@ -1680,41 +1725,6 @@ class DeepseekV4HipRadixBackend(
raise NotImplementedError("ragged attention") 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( def expand_extend_with_same_length(
self, self,
bs: int, bs: int,
@@ -11,12 +11,19 @@ Registry: nightly-amd-8-gpu-mi35x-deepseek-v4-pro-dspark suite
import os import os
import unittest import unittest
from types import SimpleNamespace from types import SimpleNamespace
from unittest import mock
import requests import requests
import torch 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.attention.dsv4.unified_kv_kernels import runtime
from sglang.kernels.ops.speculative.dspark import dspark_verify_window 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.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_amd_ci 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 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): 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): def test_build_unified_commit_inject_layout(self):
stride, ring_stride = 7, 128 stride, ring_stride = 7, 128
req_pool_indices = torch.tensor([3, 0, 5, 1], device=DEVICE, dtype=torch.int32) req_pool_indices = torch.tensor([3, 0, 5, 1], device=DEVICE, dtype=torch.int32)