[AMD] Refactor unified_kv attention metadata to data class and fuse c4/128 out_loc (#28275)
This commit is contained in:
@@ -92,6 +92,59 @@ def _create_dummy_paged_compress_data(compress_ratio: int):
|
|||||||
return None
|
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
|
@dataclass
|
||||||
class DSV4AttnMetadata:
|
class DSV4AttnMetadata:
|
||||||
page_size: int
|
page_size: int
|
||||||
@@ -122,22 +175,8 @@ class DSV4AttnMetadata:
|
|||||||
c128_topk_lengths_clamp1: Optional[torch.Tensor] = None
|
c128_topk_lengths_clamp1: Optional[torch.Tensor] = None
|
||||||
c128_topk_lengths_raw: Optional[torch.Tensor] = None
|
c128_topk_lengths_raw: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
# unified_kv: per-forward prebuilt ragged decode index
|
# unified-kv metadata
|
||||||
# SWA ring write target (req_slot*ring + pos%ring), computed once per
|
unified: Optional[UnifiedKvMetadata] = None
|
||||||
# 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
|
|
||||||
|
|
||||||
c1_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False)
|
c1_flashmla_metadata: FlashMLASchedMeta = field(init=False, repr=False)
|
||||||
c4_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_topk_lengths_raw",
|
||||||
"c4_sparse_page_indices",
|
"c4_sparse_page_indices",
|
||||||
"c4_sparse_raw_indices",
|
"c4_sparse_raw_indices",
|
||||||
"unified_swa_indices",
|
"unified",
|
||||||
"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",
|
|
||||||
],
|
],
|
||||||
assign_fields=[
|
assign_fields=[
|
||||||
# Recomputed by the recorded init_forward_metadata_in_graph op
|
# Recomputed by the recorded init_forward_metadata_in_graph op
|
||||||
# each forward; not copied across replays.
|
# each forward; not copied across replays.
|
||||||
"swa_out_cache_loc",
|
"swa_out_cache_loc",
|
||||||
"unified_swa_loc",
|
|
||||||
"c1_flashmla_metadata",
|
"c1_flashmla_metadata",
|
||||||
"c4_flashmla_metadata",
|
"c4_flashmla_metadata",
|
||||||
"c128_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.page_table.dim() == 2
|
||||||
assert (
|
assert (
|
||||||
self.raw_out_loc.shape == self.seq_lens_casual.shape
|
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.c128_page_indices = _pad_last_dim(self.c128_page_indices)
|
||||||
self.swa_page_indices = _pad_last_dim(self.swa_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 = [
|
_CP_REINDEX_FIELDS = [
|
||||||
"seq_lens_casual",
|
"seq_lens_casual",
|
||||||
"positions_casual",
|
"positions_casual",
|
||||||
@@ -1013,13 +1048,15 @@ 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]
|
||||||
|
if core.unified is None:
|
||||||
|
core.unified = UnifiedKvMetadata()
|
||||||
(
|
(
|
||||||
core.unified_swa_indices,
|
core.unified.swa_indices,
|
||||||
core.unified_swa_indptr,
|
core.unified.swa_indptr,
|
||||||
core.unified_hca_indices,
|
core.unified.hca_indices,
|
||||||
core.unified_hca_indptr,
|
core.unified.hca_indptr,
|
||||||
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=req_pool_indices[:N],
|
||||||
positions=core.positions_casual,
|
positions=core.positions_casual,
|
||||||
@@ -1035,7 +1072,7 @@ class DeepseekV4HipRadixBackend(
|
|||||||
# 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).
|
# Decode: N tokens == N reqs, positions already aligned (no repeat).
|
||||||
req_slot = req_pool_indices[:N].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)
|
||||||
@@ -1061,11 +1098,13 @@ class DeepseekV4HipRadixBackend(
|
|||||||
bid = torch.repeat_interleave(
|
bid = torch.repeat_interleave(
|
||||||
torch.arange(bs, device=device, dtype=torch.int64), extend_seq_lens
|
torch.arange(bs, device=device, dtype=torch.int64), extend_seq_lens
|
||||||
)
|
)
|
||||||
core.unified_pf_state_slot = req_pool_indices[bid]
|
if core.unified is None:
|
||||||
core.unified_pf_chunk_start = (seq_lens - extend_seq_lens)[bid]
|
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
|
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_cu_q = cu_q_per_req[bid]
|
||||||
core.unified_pf_final_pos = (seq_lens - 1)[bid]
|
core.unified.pf_final_pos = (seq_lens - 1)[bid]
|
||||||
|
|
||||||
def _forward_unified_kv(
|
def _forward_unified_kv(
|
||||||
self,
|
self,
|
||||||
@@ -1113,15 +1152,16 @@ class DeepseekV4HipRadixBackend(
|
|||||||
ring_stride=ring_stride,
|
ring_stride=ring_stride,
|
||||||
final_pos=positions,
|
final_pos=positions,
|
||||||
)
|
)
|
||||||
|
unified_metadata = core_attn_metadata.unified
|
||||||
if compress_ratio == 0:
|
if compress_ratio == 0:
|
||||||
kv_indices = core_attn_metadata.unified_swa_indices
|
kv_indices = unified_metadata.swa_indices
|
||||||
kv_indptr = core_attn_metadata.unified_swa_indptr
|
kv_indptr = unified_metadata.swa_indptr
|
||||||
elif compress_ratio == 128:
|
elif compress_ratio == 128:
|
||||||
kv_indices = core_attn_metadata.unified_hca_indices
|
kv_indices = unified_metadata.hca_indices
|
||||||
kv_indptr = core_attn_metadata.unified_hca_indptr
|
kv_indptr = unified_metadata.hca_indptr
|
||||||
elif compress_ratio == 4:
|
elif compress_ratio == 4:
|
||||||
kv_indices = core_attn_metadata.unified_csa_indices
|
kv_indices = unified_metadata.csa_indices
|
||||||
kv_indptr = core_attn_metadata.unified_csa_indptr
|
kv_indptr = unified_metadata.csa_indptr
|
||||||
runtime.fill_compress_tail(
|
runtime.fill_compress_tail(
|
||||||
indices=kv_indices,
|
indices=kv_indices,
|
||||||
indptr=kv_indptr,
|
indptr=kv_indptr,
|
||||||
@@ -1142,10 +1182,10 @@ class DeepseekV4HipRadixBackend(
|
|||||||
)
|
)
|
||||||
|
|
||||||
# prefill / extend
|
# prefill / extend
|
||||||
state_slot = core_attn_metadata.unified_pf_state_slot
|
state_slot = core_attn_metadata.unified.pf_state_slot
|
||||||
chunk_start = core_attn_metadata.unified_pf_chunk_start
|
chunk_start = core_attn_metadata.unified.pf_chunk_start
|
||||||
cu_q = core_attn_metadata.unified_pf_cu_q
|
cu_q = core_attn_metadata.unified.pf_cu_q
|
||||||
final_pos = core_attn_metadata.unified_pf_final_pos
|
final_pos = core_attn_metadata.unified.pf_final_pos
|
||||||
|
|
||||||
kpre_i, kpre_p, kext_i, kext_p = runtime.build_prefill_indices(
|
kpre_i, kpre_p, kext_i, kext_p = runtime.build_prefill_indices(
|
||||||
compress_ratio=compress_ratio,
|
compress_ratio=compress_ratio,
|
||||||
@@ -1228,7 +1268,8 @@ class DeepseekV4HipRadixBackend(
|
|||||||
"""
|
"""
|
||||||
positions = forward_batch.positions
|
positions = forward_batch.positions
|
||||||
core = getattr(self.forward_metadata, "core_attn_metadata", None)
|
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 (
|
if (
|
||||||
cached is not None
|
cached is not None
|
||||||
and not forward_batch.forward_mode.is_idle()
|
and not forward_batch.forward_mode.is_idle()
|
||||||
@@ -1495,7 +1536,9 @@ class DeepseekV4HipRadixBackend(
|
|||||||
)
|
)
|
||||||
|
|
||||||
if need_compress:
|
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()
|
core_attn_metadata.init_flashmla_related()
|
||||||
else:
|
else:
|
||||||
core_attn_metadata.c4_sparse_topk_lengths = None
|
core_attn_metadata.c4_sparse_topk_lengths = None
|
||||||
|
|||||||
@@ -511,7 +511,10 @@ class CompressorBackendMixin:
|
|||||||
elif is_unified_kv_triton():
|
elif is_unified_kv_triton():
|
||||||
kv_cache = token_to_kv_pool.get_unified_kv(layer_id)
|
kv_cache = token_to_kv_pool.get_unified_kv(layer_id)
|
||||||
page_size = 1
|
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
|
bf16_store = True
|
||||||
else:
|
else:
|
||||||
_, _, compress_kv_pool = token_to_kv_pool.layer_mapping[layer_id]
|
_, _, compress_kv_pool = token_to_kv_pool.layer_mapping[layer_id]
|
||||||
|
|||||||
@@ -3,7 +3,13 @@ from __future__ import annotations
|
|||||||
import functools
|
import functools
|
||||||
import os
|
import os
|
||||||
|
|
||||||
|
from sglang.srt.utils import is_hip
|
||||||
|
|
||||||
|
|
||||||
@functools.lru_cache(maxsize=1)
|
@functools.lru_cache(maxsize=1)
|
||||||
def is_unified_kv_triton() -> bool:
|
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"
|
||||||
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user