[AMD] Refactor unified_kv attention metadata to data class and fuse c4/128 out_loc (#28275)

This commit is contained in:
Thomas Wang
2026-06-15 02:59:01 -07:00
committed by GitHub
parent 3b419f66da
commit da12f36629
3 changed files with 105 additions and 53 deletions
@@ -92,6 +92,59 @@ def _create_dummy_paged_compress_data(compress_ratio: int):
return None
@dataclass
class UnifiedKvMetadata:
"""
unified-kv per-forward metadata
"""
# SWA ring write target (req_slot*ring + pos%ring)
swa_loc: Optional[torch.Tensor] = None
# ragged decode index streams
swa_indices: Optional[torch.Tensor] = None
swa_indptr: Optional[torch.Tensor] = None
hca_indices: Optional[torch.Tensor] = None
hca_indptr: Optional[torch.Tensor] = None
csa_indices: Optional[torch.Tensor] = None
csa_indptr: Optional[torch.Tensor] = None
# prefill/extend per-token mapping
pf_state_slot: Optional[torch.Tensor] = None
pf_chunk_start: Optional[torch.Tensor] = None
pf_cu_q: Optional[torch.Tensor] = None
pf_final_pos: Optional[torch.Tensor] = None
# 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.
c4_out_loc: Optional[torch.Tensor] = None
c128_out_loc: Optional[torch.Tensor] = None
def copy_(self, other: UnifiedKvMetadata) -> None:
copy_metadata(
src=other,
dst=self,
check_eq_fields=[],
copy_fields=[
"swa_indices",
"swa_indptr",
"hca_indices",
"hca_indptr",
"csa_indices",
"csa_indptr",
"pf_state_slot",
"pf_chunk_start",
"pf_cu_q",
"pf_final_pos",
"c4_out_loc",
"c128_out_loc",
],
# swa_loc is recomputed each forward (recorded inside cuda graphs),
# so it is rebound rather than copied across replays.
assign_fields=["swa_loc"],
)
@dataclass
class DSV4AttnMetadata:
page_size: int
@@ -122,22 +175,8 @@ class DSV4AttnMetadata:
c128_topk_lengths_clamp1: Optional[torch.Tensor] = None
c128_topk_lengths_raw: Optional[torch.Tensor] = None
# unified_kv: per-forward prebuilt ragged decode index
# SWA ring write target (req_slot*ring + pos%ring), computed once per
# forward in _attach_unified_kv_decode_streams, read by every layer's store.
unified_swa_loc: Optional[torch.Tensor] = None
unified_swa_indices: Optional[torch.Tensor] = None
unified_swa_indptr: Optional[torch.Tensor] = None
unified_hca_indices: Optional[torch.Tensor] = None
unified_hca_indptr: Optional[torch.Tensor] = None
unified_csa_indices: Optional[torch.Tensor] = None
unified_csa_indptr: Optional[torch.Tensor] = None
# unified_kv: per-forward prefill/extend per-token mapping
unified_pf_state_slot: Optional[torch.Tensor] = None
unified_pf_chunk_start: Optional[torch.Tensor] = None
unified_pf_cu_q: Optional[torch.Tensor] = None
unified_pf_final_pos: Optional[torch.Tensor] = None
# unified-kv metadata
unified: Optional[UnifiedKvMetadata] = None
c1_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False)
c4_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False)
@@ -184,29 +223,19 @@ class DSV4AttnMetadata:
"c4_sparse_topk_lengths_raw",
"c4_sparse_page_indices",
"c4_sparse_raw_indices",
"unified_swa_indices",
"unified_swa_indptr",
"unified_hca_indices",
"unified_hca_indptr",
"unified_csa_indices",
"unified_csa_indptr",
"unified_pf_state_slot",
"unified_pf_chunk_start",
"unified_pf_cu_q",
"unified_pf_final_pos",
"unified",
],
assign_fields=[
# Recomputed by the recorded init_forward_metadata_in_graph op
# each forward; not copied across replays.
"swa_out_cache_loc",
"unified_swa_loc",
"c1_flashmla_metadata",
"c4_flashmla_metadata",
"c128_flashmla_metadata",
],
)
def init_compression_metadata(self):
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
@@ -234,6 +263,12 @@ class DSV4AttnMetadata:
self.c128_page_indices = _pad_last_dim(self.c128_page_indices)
self.swa_page_indices = _pad_last_dim(self.swa_page_indices)
if unified_swa_pages:
if self.unified is None:
self.unified = UnifiedKvMetadata()
self.unified.c4_out_loc = self.c4_out_loc + unified_swa_pages
self.unified.c128_out_loc = self.c128_out_loc + unified_swa_pages
_CP_REINDEX_FIELDS = [
"seq_lens_casual",
"positions_casual",
@@ -1013,13 +1048,15 @@ class DeepseekV4HipRadixBackend(
pool = self.token_to_kv_pool
N = core.positions_casual.shape[0]
if core.unified is None:
core.unified = UnifiedKvMetadata()
(
core.unified_swa_indices,
core.unified_swa_indptr,
core.unified_hca_indices,
core.unified_hca_indptr,
core.unified_csa_indices,
core.unified_csa_indptr,
core.unified.swa_indices,
core.unified.swa_indptr,
core.unified.hca_indices,
core.unified.hca_indptr,
core.unified.csa_indices,
core.unified.csa_indptr,
) = runtime.build_decode_streams(
state_slot=req_pool_indices[:N],
positions=core.positions_casual,
@@ -1035,7 +1072,7 @@ class DeepseekV4HipRadixBackend(
# SWA ring write target, same value for every layer this forward.
# Decode: N tokens == N reqs, positions already aligned (no repeat).
req_slot = req_pool_indices[:N].to(torch.int64)
core.unified_swa_loc = (
core.unified.swa_loc = (
req_slot * pool.unified_swa_ring_size
+ core.positions_casual.to(torch.int64) % pool.unified_swa_ring_size
).to(torch.int32)
@@ -1061,11 +1098,13 @@ class DeepseekV4HipRadixBackend(
bid = torch.repeat_interleave(
torch.arange(bs, device=device, dtype=torch.int64), extend_seq_lens
)
core.unified_pf_state_slot = req_pool_indices[bid]
core.unified_pf_chunk_start = (seq_lens - extend_seq_lens)[bid]
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]
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]
core.unified.pf_cu_q = cu_q_per_req[bid]
core.unified.pf_final_pos = (seq_lens - 1)[bid]
def _forward_unified_kv(
self,
@@ -1113,15 +1152,16 @@ class DeepseekV4HipRadixBackend(
ring_stride=ring_stride,
final_pos=positions,
)
unified_metadata = core_attn_metadata.unified
if compress_ratio == 0:
kv_indices = core_attn_metadata.unified_swa_indices
kv_indptr = core_attn_metadata.unified_swa_indptr
kv_indices = unified_metadata.swa_indices
kv_indptr = unified_metadata.swa_indptr
elif compress_ratio == 128:
kv_indices = core_attn_metadata.unified_hca_indices
kv_indptr = core_attn_metadata.unified_hca_indptr
kv_indices = unified_metadata.hca_indices
kv_indptr = unified_metadata.hca_indptr
elif compress_ratio == 4:
kv_indices = core_attn_metadata.unified_csa_indices
kv_indptr = core_attn_metadata.unified_csa_indptr
kv_indices = unified_metadata.csa_indices
kv_indptr = unified_metadata.csa_indptr
runtime.fill_compress_tail(
indices=kv_indices,
indptr=kv_indptr,
@@ -1142,10 +1182,10 @@ class DeepseekV4HipRadixBackend(
)
# prefill / extend
state_slot = core_attn_metadata.unified_pf_state_slot
chunk_start = core_attn_metadata.unified_pf_chunk_start
cu_q = core_attn_metadata.unified_pf_cu_q
final_pos = core_attn_metadata.unified_pf_final_pos
state_slot = core_attn_metadata.unified.pf_state_slot
chunk_start = core_attn_metadata.unified.pf_chunk_start
cu_q = core_attn_metadata.unified.pf_cu_q
final_pos = core_attn_metadata.unified.pf_final_pos
kpre_i, kpre_p, kext_i, kext_p = runtime.build_prefill_indices(
compress_ratio=compress_ratio,
@@ -1228,7 +1268,8 @@ class DeepseekV4HipRadixBackend(
"""
positions = forward_batch.positions
core = getattr(self.forward_metadata, "core_attn_metadata", None)
cached = core.unified_swa_loc if core is not None else None
unified = getattr(core, "unified", None) if core is not None else None
cached = unified.swa_loc if unified is not None else None
if (
cached is not None
and not forward_batch.forward_mode.is_idle()
@@ -1495,7 +1536,9 @@ class DeepseekV4HipRadixBackend(
)
if need_compress:
core_attn_metadata.init_compression_metadata()
core_attn_metadata.init_compression_metadata(
unified_swa_pages=getattr(self.token_to_kv_pool, "unified_swa_pages", 0)
)
core_attn_metadata.init_flashmla_related()
else:
core_attn_metadata.c4_sparse_topk_lengths = None
@@ -511,7 +511,10 @@ class CompressorBackendMixin:
elif is_unified_kv_triton():
kv_cache = token_to_kv_pool.get_unified_kv(layer_id)
page_size = 1
out_loc = out_loc + token_to_kv_pool.unified_swa_pages
out_loc = getattr(
self.forward_metadata.core_metadata.unified,
f"c{compressor.ratio}_out_loc",
)
bf16_store = True
else:
_, _, compress_kv_pool = token_to_kv_pool.layer_mapping[layer_id]
@@ -3,7 +3,13 @@ from __future__ import annotations
import functools
import os
from sglang.srt.utils import is_hip
@functools.lru_cache(maxsize=1)
def is_unified_kv_triton() -> bool:
return os.environ.get("SGLANG_HACK_FLASHMLA_BACKEND", "") == "unified_kv_triton"
# unified_kv_triton is only implemented on HIP (ROCm)
return (
is_hip()
and os.environ.get("SGLANG_HACK_FLASHMLA_BACKEND", "") == "unified_kv_triton"
)