From da12f36629a1774bfab156f17ce87faa93a4e1e7 Mon Sep 17 00:00:00 2001 From: Thomas Wang Date: Mon, 15 Jun 2026 17:59:01 +0800 Subject: [PATCH] [AMD] Refactor unified_kv attention metadata to data class and fuse c4/128 out_loc (#28275) --- .../deepseek_v4_backend_hip_radix.py | 145 ++++++++++++------ .../layers/attention/dsv4/compressor_v2.py | 5 +- .../dsv4/unified_kv_kernels/env_gate.py | 8 +- 3 files changed, 105 insertions(+), 53 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 6602a5a62..91d16a20e 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 @@ -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 diff --git a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py index c3caaf86f..3d5c4b16f 100644 --- a/python/sglang/srt/layers/attention/dsv4/compressor_v2.py +++ b/python/sglang/srt/layers/attention/dsv4/compressor_v2.py @@ -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] diff --git a/python/sglang/srt/layers/attention/dsv4/unified_kv_kernels/env_gate.py b/python/sglang/srt/layers/attention/dsv4/unified_kv_kernels/env_gate.py index 1835798d2..c6636f9dc 100644 --- a/python/sglang/srt/layers/attention/dsv4/unified_kv_kernels/env_gate.py +++ b/python/sglang/srt/layers/attention/dsv4/unified_kv_kernels/env_gate.py @@ -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" + )