Facade DSA index-cache: MTP topk-reuse state + index-K storage (#28609)
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user