[AMD] Fix Dspark accept length and reduce host bubble on DSV4 (#39116)
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user