[AMD] Run V4 MTP target-verify through the decode kernel (#34597)
Co-authored-by: RolaoDenthu <xinyis10@illinois.edu> Co-authored-by: 1am9trash <1am9trash> Co-authored-by: kk <43161300+kkHuang-amd@users.noreply.github.com>
This commit is contained in:
co-authored by
RolaoDenthu
1am9trash
kk
parent
fd1e04d952
commit
34206c0017
@@ -105,6 +105,12 @@ class UnifiedKvMetadata:
|
|||||||
pf_cu_q: Optional[torch.Tensor] = None
|
pf_cu_q: Optional[torch.Tensor] = None
|
||||||
pf_final_pos: Optional[torch.Tensor] = None
|
pf_final_pos: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
|
# Per-token req-slot map used by the SWA ring store, precomputed once per
|
||||||
|
# step so the forward store does not recompute a repeat_interleave per layer.
|
||||||
|
# Read by the target-verify store (num_draft*bs tokens); for plain decode it
|
||||||
|
# equals req_pool_indices and is unused (the decode store reads that live).
|
||||||
|
verify_store_state_slot: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
# SWA-page-offset compressed-store locations (= c*_out_loc + unified_swa_pages),
|
# SWA-page-offset compressed-store locations (= c*_out_loc + unified_swa_pages),
|
||||||
# precomputed once per step to drop the per-layer int add in the store path.
|
# precomputed once per step to drop the per-layer int add in the store path.
|
||||||
c4_out_loc: Optional[torch.Tensor] = None
|
c4_out_loc: Optional[torch.Tensor] = None
|
||||||
@@ -126,6 +132,7 @@ class UnifiedKvMetadata:
|
|||||||
"pf_chunk_start",
|
"pf_chunk_start",
|
||||||
"pf_cu_q",
|
"pf_cu_q",
|
||||||
"pf_final_pos",
|
"pf_final_pos",
|
||||||
|
"verify_store_state_slot",
|
||||||
"c4_out_loc",
|
"c4_out_loc",
|
||||||
"c128_out_loc",
|
"c128_out_loc",
|
||||||
],
|
],
|
||||||
@@ -547,6 +554,7 @@ class DeepseekV4HipRadixBackend(
|
|||||||
use_prefill_cuda_graph: bool = False,
|
use_prefill_cuda_graph: bool = False,
|
||||||
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,
|
||||||
) -> DSV4Metadata:
|
) -> DSV4Metadata:
|
||||||
if extend_start_loc is not None:
|
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 (
|
||||||
@@ -590,6 +598,14 @@ class DeepseekV4HipRadixBackend(
|
|||||||
num_tokens,
|
num_tokens,
|
||||||
need_compress=need_compress,
|
need_compress=need_compress,
|
||||||
)
|
)
|
||||||
|
if attach_decode_streams:
|
||||||
|
# Target-verify runs through the unified_kv DECODE kernel, so build
|
||||||
|
# per-token decode streams here. req_pool_indices_repeated is the
|
||||||
|
# per-token (num_draft*bs -> bs) req-slot map produced by the prefill
|
||||||
|
# expansion above.
|
||||||
|
self._attach_unified_kv_decode_streams(
|
||||||
|
core_attn_metadata, req_pool_indices_repeated
|
||||||
|
)
|
||||||
indexer_metadata = (
|
indexer_metadata = (
|
||||||
self.init_forward_metadata_indexer(core_attn_metadata)
|
self.init_forward_metadata_indexer(core_attn_metadata)
|
||||||
if need_compress
|
if need_compress
|
||||||
@@ -712,6 +728,7 @@ class DeepseekV4HipRadixBackend(
|
|||||||
use_prefill_cuda_graph=use_prefill_cuda_graph,
|
use_prefill_cuda_graph=use_prefill_cuda_graph,
|
||||||
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,
|
||||||
)
|
)
|
||||||
|
|
||||||
def make_forward_metadata_from_raw_verify(
|
def make_forward_metadata_from_raw_verify(
|
||||||
@@ -1133,9 +1150,14 @@ class DeepseekV4HipRadixBackend(
|
|||||||
self.forward_metadata = current_raw
|
self.forward_metadata = current_raw
|
||||||
|
|
||||||
def _attach_unified_kv_decode_streams(
|
def _attach_unified_kv_decode_streams(
|
||||||
self, core: DSV4AttnMetadata, req_pool_indices: torch.Tensor
|
self, core: DSV4AttnMetadata, state_slot: torch.Tensor
|
||||||
) -> None:
|
) -> None:
|
||||||
"""build the ragged decode index streams once per forward"""
|
"""build the ragged decode index streams once per forward.
|
||||||
|
|
||||||
|
``state_slot`` is the per-row req-slot map: decode passes
|
||||||
|
``req_pool_indices`` (1 token per req), target-verify passes
|
||||||
|
``req_pool_indices_repeated`` (the per-token num_draft*bs -> bs map) so
|
||||||
|
the same builder produces per-draft-token decode streams."""
|
||||||
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,
|
||||||
)
|
)
|
||||||
@@ -1146,6 +1168,7 @@ class DeepseekV4HipRadixBackend(
|
|||||||
|
|
||||||
pool = self.token_to_kv_pool
|
pool = self.token_to_kv_pool
|
||||||
N = core.positions_casual.shape[0]
|
N = core.positions_casual.shape[0]
|
||||||
|
state_slot = state_slot[:N]
|
||||||
if core.unified is None:
|
if core.unified is None:
|
||||||
core.unified = UnifiedKvMetadata()
|
core.unified = UnifiedKvMetadata()
|
||||||
(
|
(
|
||||||
@@ -1156,7 +1179,7 @@ class DeepseekV4HipRadixBackend(
|
|||||||
core.unified.csa_indices,
|
core.unified.csa_indices,
|
||||||
core.unified.csa_indptr,
|
core.unified.csa_indptr,
|
||||||
) = runtime.build_decode_streams(
|
) = runtime.build_decode_streams(
|
||||||
state_slot=req_pool_indices[:N],
|
state_slot=state_slot,
|
||||||
positions=core.positions_casual,
|
positions=core.positions_casual,
|
||||||
swa_len=core.swa_topk_lengths,
|
swa_len=core.swa_topk_lengths,
|
||||||
hca_len=core.c128_topk_lengths_raw,
|
hca_len=core.c128_topk_lengths_raw,
|
||||||
@@ -1168,12 +1191,15 @@ class DeepseekV4HipRadixBackend(
|
|||||||
swa_pages=pool.unified_swa_pages,
|
swa_pages=pool.unified_swa_pages,
|
||||||
)
|
)
|
||||||
# SWA ring write target, same value for every layer this forward.
|
# SWA ring write target, same value for every layer this forward.
|
||||||
# Decode: N tokens == N reqs, positions already aligned (no repeat).
|
req_slot = state_slot.to(torch.int64)
|
||||||
req_slot = req_pool_indices[:N].to(torch.int64)
|
|
||||||
core.unified.swa_loc = (
|
core.unified.swa_loc = (
|
||||||
req_slot * pool.unified_swa_ring_size
|
req_slot * pool.unified_swa_ring_size
|
||||||
+ core.positions_casual.to(torch.int64) % pool.unified_swa_ring_size
|
+ core.positions_casual.to(torch.int64) % pool.unified_swa_ring_size
|
||||||
).to(torch.int32)
|
).to(torch.int32)
|
||||||
|
# Per-token req-slot map for the SWA ring store, read directly by the
|
||||||
|
# forward store (target-verify) instead of recomputing a repeat_interleave
|
||||||
|
# per layer. Harmless for plain decode (its store reads req_pool_indices).
|
||||||
|
core.unified.verify_store_state_slot = state_slot
|
||||||
|
|
||||||
def _attach_unified_kv_prefill_meta(
|
def _attach_unified_kv_prefill_meta(
|
||||||
self,
|
self,
|
||||||
@@ -1248,9 +1274,18 @@ class DeepseekV4HipRadixBackend(
|
|||||||
c128_pi = getattr(core_attn_metadata, "c128_page_indices", None)
|
c128_pi = getattr(core_attn_metadata, "c128_page_indices", None)
|
||||||
c4_pi = getattr(core_attn_metadata, "c4_sparse_page_indices", None)
|
c4_pi = getattr(core_attn_metadata, "c4_sparse_page_indices", None)
|
||||||
|
|
||||||
# decode
|
# Target-verify runs through the unified_kv DECODE kernel, same path as
|
||||||
is_decode = forward_batch.forward_mode.is_decode_or_idle()
|
# decode; its per-token decode streams were built in metadata.
|
||||||
|
verify_as_decode = forward_batch.forward_mode.is_target_verify()
|
||||||
|
is_decode = forward_batch.forward_mode.is_decode_or_idle() or verify_as_decode
|
||||||
if is_decode:
|
if is_decode:
|
||||||
|
if verify_as_decode:
|
||||||
|
# Per-token (num_draft*bs -> bs) req-slot map, precomputed once
|
||||||
|
# per step in _attach_unified_kv_decode_streams. Writing every
|
||||||
|
# draft token's K into the ring is safe: spec_extra room prevents
|
||||||
|
# clobbering the window history same-step tokens still read.
|
||||||
|
state_slot = core_attn_metadata.unified.verify_store_state_slot[:T]
|
||||||
|
else:
|
||||||
state_slot = forward_batch.req_pool_indices[:T]
|
state_slot = forward_batch.req_pool_indices[:T]
|
||||||
if save_kv_cache:
|
if save_kv_cache:
|
||||||
runtime.store_swa_into_unified(
|
runtime.store_swa_into_unified(
|
||||||
|
|||||||
Reference in New Issue
Block a user