Facade DSA index-cache: MTP topk-reuse state + index-K storage (#28609)

This commit is contained in:
Xinyuan Tong
2026-08-06 00:34:31 -07:00
committed by GitHub
parent 735995e7bd
commit 31c1e5943f
9 changed files with 672 additions and 388 deletions
+20 -125
View File
@@ -38,7 +38,6 @@ import torch
import triton
import triton.language as tl
from sglang.kernels.ops.attention.dsa import index_buf_accessor
from sglang.kernels.ops.attention.dsa.quant_k_cache import (
quantize_k_cache,
quantize_k_cache_separate,
@@ -59,6 +58,7 @@ from sglang.srt.layers.quantization.fp4_kv_cache_quant_method import (
)
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.allocator.mamba import MambaSlotAllocator
from sglang.srt.mem_cache.index_key_cache import IndexKeyCache
from sglang.srt.mem_cache.kv_vmm_backing import KvVmmBufferOwner
from sglang.srt.mem_cache.layout.page_major import (
build_page_major_mamba_views,
@@ -4378,58 +4378,28 @@ class DSATokenToKVPool(MLATokenToKVPool):
), f"HIP legacy DSA path requires page_size == 1, got {self.page_size}"
else:
assert self.page_size == 64
self._create_index_buffers()
self.index_key_cache = self._create_index_key_cache()
self._finalize_allocation_log(size)
def _index_buffer_shape(self, num_pages: int) -> tuple[int, int]:
return (
num_pages,
self.page_size
* (self.index_head_dim + self.index_head_dim // self.quant_block_size * 4),
)
def _create_index_key_cache(self) -> IndexKeyCache:
return IndexKeyCache(self, self.index_buf_size)
def _create_index_buffers(self):
num_pages = (self.index_buf_size + self.page_size + 1) // self.page_size
with (
torch.cuda.use_mem_pool(self.custom_mem_pool)
if self.custom_mem_pool
else nullcontext()
):
self.index_k_with_scale_buffer = [
torch.zeros(
# Layout:
# ref: test_attention.py :: kv_cache_cast_to_fp8
# shape: (num_pages, page_size 64 * head_dim 128 + page_size 64 * fp32_nbytes 4)
# data: for page i,
# * buf[i, :page_size * head_dim] for fp8 data
# * buf[i, page_size * head_dim:].view(float32) for scale
self._index_buffer_shape(num_pages),
dtype=self.index_k_with_scale_buffer_dtype,
device=self.device,
)
for _ in range(self.layer_num)
]
@property
def index_k_with_scale_buffer(self):
# Preserve direct HiCache access while storage lives behind the facade.
return self.index_key_cache.buffer
def _clear_buffers(self):
super()._clear_buffers()
del self.index_k_with_scale_buffer
self.index_key_cache.clear()
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
"""Move latent KV and the DSA indexer cache (key + scale) in lockstep."""
super().move_kv_cache(tgt_loc, src_loc)
if tgt_loc.numel() == 0:
return
tgt_loc_flat = tgt_loc.view(-1).long()
src_loc_flat = src_loc.view(-1).long()
for index_k in self.index_k_with_scale_buffer:
index_k[tgt_loc_flat] = index_k[src_loc_flat]
self.index_key_cache.move(tgt_loc, src_loc)
def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor:
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
return self.index_k_with_scale_buffer[layer_id - self.start_layer]
return self.index_key_cache.get_local_buffer(layer_id)
def get_index_k_continuous(
self,
@@ -4437,12 +4407,7 @@ class DSATokenToKVPool(MLATokenToKVPool):
seq_len: int,
page_indices: torch.Tensor,
):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
buf = self.index_k_with_scale_buffer[layer_id - self.start_layer]
return index_buf_accessor.GetK.execute(
self, buf, seq_len=seq_len, page_indices=page_indices
)
return self.index_key_cache.get_k_continuous(layer_id, seq_len, page_indices)
def get_index_k_scale_continuous(
self,
@@ -4450,11 +4415,8 @@ class DSATokenToKVPool(MLATokenToKVPool):
seq_len: int,
page_indices: torch.Tensor,
):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
buf = self.index_k_with_scale_buffer[layer_id - self.start_layer]
return index_buf_accessor.GetS.execute(
self, buf, seq_len=seq_len, page_indices=page_indices
return self.index_key_cache.get_k_scale_continuous(
layer_id, seq_len, page_indices
)
def get_index_k_scale_buffer(
@@ -4465,27 +4427,8 @@ class DSATokenToKVPool(MLATokenToKVPool):
seq_len_sum: int,
max_seq_len: int,
):
"""
Fused method to get both index K and scale data in a single call using Triton.
More efficient than calling get_index_k_continuous and get_index_k_scale_continuous separately.
:param layer_id: Layer index
:param seq_len: Sequence length
:param page_indices: Page indices tensor
:return: tuple of (k_fp8, k_scale) where
k_fp8: (seq_len, index_head_dim), uint8
k_scale: (seq_len, 4), uint8
"""
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
buf = self.index_k_with_scale_buffer[layer_id - self.start_layer]
return index_buf_accessor.GetKAndS.execute(
self,
buf,
page_indices=page_indices,
seq_len_tensor=seq_len_tensor,
seq_len_sum=seq_len_sum,
max_seq_len=max_seq_len,
return self.index_key_cache.get_k_and_scale(
layer_id, seq_len_tensor, page_indices, seq_len_sum, max_seq_len
)
def set_index_k_scale_buffer(
@@ -4495,68 +4438,20 @@ class DSATokenToKVPool(MLATokenToKVPool):
index_k: torch.Tensor,
index_k_scale: torch.Tensor,
) -> None:
buf = self.index_k_with_scale_buffer[layer_id - self.start_layer]
index_buf_accessor.SetKAndS.execute(
pool=self, buf=buf, loc=loc, index_k=index_k, index_k_scale=index_k_scale
)
self.index_key_cache.store_quantized(layer_id, loc, index_k, index_k_scale)
def get_cpu_copy(self, indices, mamba_indices=None):
# DSA keeps a page-indexed index_k_with_scale_buffer alongside kv_buffer.
# Retract frees the slots/pages and they get reused by other reqs'
# set_index_k_scale_buffer, so we must offload it here too -- otherwise
# resume restores kv_buffer but leaves foreign index/scale in place and
# DSA attention reads garbage at those token positions.
kv_cache_cpu = super().get_cpu_copy(indices, mamba_indices=mamba_indices)
page_indices = indices[:: self.page_size] // self.page_size
torch.cuda.synchronize()
index_k_cpu = []
chunk_size = self.cpu_offloading_chunk_size
page_chunk_size = max(1, chunk_size // self.page_size)
for layer_id in range(self.layer_num):
index_k_cpu.append([])
for i in range(0, len(page_indices), page_chunk_size):
chunk_page_indices = page_indices[i : i + page_chunk_size]
idx_cpu = self.index_k_with_scale_buffer[layer_id][
chunk_page_indices
].to("cpu", non_blocking=True)
index_k_cpu[-1].append(idx_cpu)
torch.cuda.synchronize()
return {"kv": kv_cache_cpu, "index_k": index_k_cpu}
return {"kv": kv_cache_cpu, "index_k": self.index_key_cache.cpu_copy(indices)}
def load_cpu_copy(self, kv_cache_cpu_dict, indices, mamba_indices=None):
super().load_cpu_copy(
kv_cache_cpu_dict["kv"], indices, mamba_indices=mamba_indices
)
page_indices = indices[:: self.page_size] // self.page_size
index_k_cpu = kv_cache_cpu_dict["index_k"]
torch.cuda.synchronize()
chunk_size = self.cpu_offloading_chunk_size
page_chunk_size = max(1, chunk_size // self.page_size)
for layer_id in range(self.layer_num):
for i in range(0, len(page_indices), page_chunk_size):
chunk_page_indices = page_indices[i : i + page_chunk_size]
idx_cpu = index_k_cpu[layer_id][i // page_chunk_size]
assert idx_cpu.shape[0] == len(chunk_page_indices)
idx_chunk = idx_cpu.to(
self.index_k_with_scale_buffer[0].device, non_blocking=True
)
self.index_k_with_scale_buffer[layer_id][chunk_page_indices] = idx_chunk
torch.cuda.synchronize()
self.index_key_cache.load_cpu_copy(kv_cache_cpu_dict["index_k"], indices)
def get_state_buf_infos(self):
data_ptrs = [
self.index_k_with_scale_buffer[i].data_ptr() for i in range(self.layer_num)
]
data_lens = [
self.index_k_with_scale_buffer[i].nbytes for i in range(self.layer_num)
]
item_lens = [
self.index_k_with_scale_buffer[i][0].nbytes for i in range(self.layer_num)
]
return data_ptrs, data_lens, item_lens
return self.index_key_cache.state_buf_infos()
def get_kv_size_bytes(self):
kv_size_bytes = super().get_kv_size_bytes()