[HiCache]: Support DeepSeek v32 cpu offloading (#17415)

Co-authored-by: 晟海 <huangtingwei.htw@antgroup.com>
This commit is contained in:
zhangheng
2026-02-03 10:07:37 +08:00
committed by GitHub
parent a1bbc892af
commit 180594358b
3 changed files with 334 additions and 32 deletions

View File

@@ -21,10 +21,15 @@ from sglang.srt.mem_cache.base_prefix_cache import (
MatchPrefixParams,
MatchResult,
)
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, MLATokenToKVPool
from sglang.srt.mem_cache.memory_pool import (
MHATokenToKVPool,
MLATokenToKVPool,
NSATokenToKVPool,
)
from sglang.srt.mem_cache.memory_pool_host import (
MHATokenToKVPoolHost,
MLATokenToKVPoolHost,
NSATokenToKVPoolHost,
)
from sglang.srt.mem_cache.radix_cache import (
RadixCache,
@@ -70,6 +75,15 @@ class HiRadixCache(RadixCache):
server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
)
elif isinstance(self.kv_cache, NSATokenToKVPool):
self.token_to_kv_pool_host = NSATokenToKVPoolHost(
self.kv_cache,
server_args.hicache_ratio,
server_args.hicache_size,
self.page_size,
server_args.hicache_mem_layout,
allocator_type=server_args.hicache_storage_backend,
)
elif isinstance(self.kv_cache, MLATokenToKVPool):
self.token_to_kv_pool_host = MLATokenToKVPoolHost(
self.kv_cache,

View File

@@ -15,7 +15,12 @@ from sglang.jit_kernel.hicache import (
from sglang.jit_kernel.hicache import (
transfer_hicache_one_layer as jit_transfer_hicache_one_layer,
)
from sglang.srt.mem_cache.memory_pool import KVCache, MHATokenToKVPool, MLATokenToKVPool
from sglang.srt.mem_cache.memory_pool import (
KVCache,
MHATokenToKVPool,
MLATokenToKVPool,
NSATokenToKVPool,
)
from sglang.srt.utils import is_cuda, is_npu, is_xpu
_is_cuda = is_cuda()
@@ -689,7 +694,9 @@ class MLATokenToKVPoolHost(HostKVCache):
pin_memory: bool = True,
device: str = "cpu",
allocator_type: str = "default",
override_kv_cache_dim: Optional[int] = None,
):
self.override_kv_cache_dim = override_kv_cache_dim
super().__init__(
device_pool,
host_to_device_ratio,
@@ -711,13 +718,10 @@ class MLATokenToKVPoolHost(HostKVCache):
self.kv_lora_rank = self.device_pool.kv_lora_rank
self.qk_rope_head_dim = self.device_pool.qk_rope_head_dim
self.layer_num = self.device_pool.layer_num
return (
(self.kv_lora_rank + self.qk_rope_head_dim)
* 1
* self.dtype.itemsize
* self.layer_num
self.kv_cache_dim = self.override_kv_cache_dim or (
self.kv_lora_rank + self.qk_rope_head_dim
)
return self.kv_cache_dim * self.dtype.itemsize * self.layer_num
def get_ksize_per_token(self):
return self.get_size_per_token()
@@ -728,14 +732,14 @@ class MLATokenToKVPoolHost(HostKVCache):
self.layer_num,
self.size,
1,
self.kv_lora_rank + self.qk_rope_head_dim,
self.kv_cache_dim,
)
elif self.layout == "page_first":
dims = (
self.size,
self.layer_num,
1,
self.kv_lora_rank + self.qk_rope_head_dim,
self.kv_cache_dim,
)
elif self.layout == "page_first_direct":
dims = (
@@ -743,7 +747,7 @@ class MLATokenToKVPoolHost(HostKVCache):
self.layer_num,
self.page_size,
1,
self.kv_lora_rank + self.qk_rope_head_dim,
self.kv_cache_dim,
)
# Ascend-specific: Aligns with NPUMLATokenToKVPool layout
# Separately allocate k_buffer and v_buffer for easier data transfer.
@@ -783,9 +787,7 @@ class MLATokenToKVPoolHost(HostKVCache):
return self.k_buffer
else:
raise ValueError(f"Unsupported layout: {self.layout}")
self.token_stride_size = (
self.kv_lora_rank + self.qk_rope_head_dim
) * self.dtype.itemsize
self.token_stride_size = self.kv_cache_dim * self.dtype.itemsize
self.layout_dim = self.token_stride_size * self.layer_num
alloc_func = ALLOC_MEMORY_FUNCS[self.device_pool.device]
@@ -946,7 +948,7 @@ class MLATokenToKVPoolHost(HostKVCache):
self.layer_num,
self.page_size,
1,
self.kv_lora_rank + self.qk_rope_head_dim,
self.kv_cache_dim,
),
dtype=self.dtype,
device=self.device,
@@ -959,14 +961,14 @@ class MLATokenToKVPoolHost(HostKVCache):
self.layer_num,
self.page_size,
1,
self.kv_lora_rank + self.qk_rope_head_dim,
self.kv_cache_dim,
)
elif self.layout == "page_first":
self.kv_buffer[index : index + self.page_size, :, :, :] = data_page.reshape(
self.page_size,
self.layer_num,
1,
self.kv_lora_rank + self.qk_rope_head_dim,
self.kv_cache_dim,
)
elif self.layout == "page_first_direct":
real_index = index // self.page_size
@@ -975,7 +977,7 @@ class MLATokenToKVPoolHost(HostKVCache):
self.layer_num,
self.page_size,
1,
self.kv_lora_rank + self.qk_rope_head_dim,
self.kv_cache_dim,
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
@@ -993,20 +995,11 @@ class MLATokenToKVPoolHost(HostKVCache):
for layer_id in range(self.layer_num):
k_ptr = (
kv_buffer_data_ptr
+ indices[index]
* (self.kv_lora_rank + self.qk_rope_head_dim)
* self.dtype.itemsize
+ layer_id
* self.size
* (self.kv_lora_rank + self.qk_rope_head_dim)
* self.dtype.itemsize
+ indices[index] * self.kv_cache_dim * self.dtype.itemsize
+ layer_id * self.size * self.kv_cache_dim * self.dtype.itemsize
)
ptr_list.append(k_ptr)
element_size = (
self.dtype.itemsize
* self.page_size
* (self.kv_lora_rank + self.qk_rope_head_dim)
)
element_size = self.dtype.itemsize * self.page_size * self.kv_cache_dim
element_size_list = [element_size] * len(ptr_list)
elif self.layout in ["page_first", "page_first_direct"]:
for index in range(0, len(indices), self.page_size):
@@ -1014,7 +1007,7 @@ class MLATokenToKVPoolHost(HostKVCache):
kv_buffer_data_ptr
+ indices[index]
* self.layer_num
* (self.kv_lora_rank + self.qk_rope_head_dim)
* self.kv_cache_dim
* self.dtype.itemsize
)
ptr_list.append(k_ptr)
@@ -1022,9 +1015,174 @@ class MLATokenToKVPoolHost(HostKVCache):
self.layer_num
* self.dtype.itemsize
* self.page_size
* (self.kv_lora_rank + self.qk_rope_head_dim)
* self.kv_cache_dim
)
element_size_list = [element_size] * len(ptr_list)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
return ptr_list, element_size_list
class NSATokenToKVPoolHost(MLATokenToKVPoolHost):
device_pool: NSATokenToKVPool
def __init__(
self,
device_pool: NSATokenToKVPool,
host_to_device_ratio: float,
host_size: int,
page_size: int,
layout: str,
pin_memory: bool = True,
device: str = "cpu",
allocator_type: str = "default",
):
# Initialize indexer metadata before HostKVCache.__init__ calls get_size_per_token.
self.index_head_dim = device_pool.index_head_dim
self.indexer_quant_block_size = device_pool.quant_block_size
self.indexer_dtype = NSATokenToKVPool.index_k_with_scale_buffer_dtype
self.indexer_size_per_token = (
self.index_head_dim
+ self.index_head_dim // self.indexer_quant_block_size * 4
)
super().__init__(
device_pool,
host_to_device_ratio,
host_size,
page_size,
layout,
pin_memory,
device,
allocator_type,
override_kv_cache_dim=device_pool.kv_cache_dim,
)
self.indexer_page_stride_size = (
self.indexer_size_per_token * self.page_size * self.indexer_dtype.itemsize
)
self.indexer_page_num = (self.size + self.page_size + 1) // self.page_size
self._init_indexer_buffers()
logger.info(
f"NSATokenToKVPoolHost initialized with indexer page stride size: {self.indexer_page_stride_size}, page num: {self.indexer_page_num}"
)
def get_size_per_token(self):
base = super().get_size_per_token()
return (
base
+ self.indexer_size_per_token * self.layer_num * self.indexer_dtype.itemsize
)
def _init_indexer_buffers(self):
alloc_func = ALLOC_MEMORY_FUNCS[self.device_pool.device]
self.index_k_with_scale_buffer = [
alloc_func(
(self.indexer_page_num, self.indexer_page_stride_size),
dtype=self.indexer_dtype,
device=self.device,
pin_memory=self.pin_memory,
allocator=self.allocator,
)
for _ in range(self.layer_num)
]
self.index_k_data_refs = [
self.index_k_with_scale_buffer[i] for i in range(self.layer_num)
]
self.index_k_data_ptrs = torch.tensor(
[x.data_ptr() for x in self.index_k_data_refs],
dtype=torch.uint64,
device=self.device_pool.device,
)
self.index_k_device_ptrs = torch.tensor(
[x.data_ptr() for x in self.device_pool.index_k_with_scale_buffer],
dtype=torch.uint64,
device=self.device_pool.device,
)
def _get_indexer_page_indices(self, host_indices, device_indices):
if host_indices.numel() == 0:
return host_indices, device_indices
if host_indices.numel() % self.page_size != 0:
raise ValueError(
"Index buffer transfer expects page-aligned indices for NSA."
)
host_page_indices = (
host_indices.reshape(-1, self.page_size)[:, 0] // self.page_size
)
device_page_indices = (
device_indices.reshape(-1, self.page_size)[:, 0] // self.page_size
)
return host_page_indices, device_page_indices
def _load_indexer_to_device_per_layer(
self, device_pool, host_indices, device_indices, layer_id, io_backend
):
host_page_indices, device_page_indices = self._get_indexer_page_indices(
host_indices, device_indices
)
use_kernel = io_backend == "kernel" and self.indexer_page_stride_size % 8 == 0
if use_kernel:
transfer_kv_per_layer_mla(
src=self.index_k_with_scale_buffer[layer_id],
dst=device_pool.index_k_with_scale_buffer[layer_id],
src_indices=host_page_indices,
dst_indices=device_page_indices,
item_size=self.indexer_page_stride_size,
)
else:
transfer_kv_direct(
src_layers=[self.index_k_with_scale_buffer[layer_id]],
dst_layers=[device_pool.index_k_with_scale_buffer[layer_id]],
src_indices=host_page_indices,
dst_indices=device_page_indices,
page_size=1,
)
def _backup_indexer_from_device_all_layer(
self, device_pool, host_indices, device_indices, io_backend
):
host_page_indices, device_page_indices = self._get_indexer_page_indices(
host_indices, device_indices
)
use_kernel = io_backend == "kernel" and self.indexer_page_stride_size % 8 == 0
if use_kernel:
transfer_kv_all_layer_mla(
src_layers=self.index_k_device_ptrs,
dst_layers=self.index_k_data_ptrs,
src_indices=device_page_indices,
dst_indices=host_page_indices,
item_size=self.indexer_page_stride_size,
num_layers=self.layer_num,
)
else:
transfer_kv_direct(
src_layers=device_pool.index_k_with_scale_buffer,
dst_layers=self.index_k_with_scale_buffer,
src_indices=device_page_indices,
dst_indices=host_page_indices,
page_size=1,
)
def load_to_device_per_layer(
self,
device_pool,
host_indices,
device_indices,
layer_id,
io_backend,
):
super().load_to_device_per_layer(
device_pool, host_indices, device_indices, layer_id, io_backend
)
self._load_indexer_to_device_per_layer(
device_pool, host_indices, device_indices, layer_id, io_backend
)
def backup_from_device_all_layer(
self, device_pool, host_indices, device_indices, io_backend
):
super().backup_from_device_all_layer(
device_pool, host_indices, device_indices, io_backend
)
self._backup_indexer_from_device_all_layer(
device_pool, host_indices, device_indices, io_backend
)

View File

@@ -0,0 +1,130 @@
import unittest
import torch
from sglang.srt.mem_cache.memory_pool import NSATokenToKVPool
from sglang.srt.mem_cache.memory_pool_host import (
ALLOC_MEMORY_FUNCS,
NSATokenToKVPoolHost,
alloc_with_pin_memory,
)
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=3, suite="stage-b-test-small-1-gpu")
class TestNSAHiCacheTransfer(unittest.TestCase):
def setUp(self):
if not torch.cuda.is_available():
self.skipTest("CUDA is required for NSA host transfer tests.")
if is_npu() or is_xpu():
self.skipTest("NSA host transfer tests only support CUDA/ROCm.")
if not (is_cuda() or is_hip()):
self.skipTest("CUDA/ROCm not available.")
@staticmethod
def _token_indices_for_pages(pages: torch.Tensor, page_size: int, device: str):
parts = [
torch.arange(
int(page_id) * page_size,
(int(page_id) + 1) * page_size,
device=device,
dtype=torch.int64,
)
for page_id in pages.tolist()
]
return torch.cat(parts, dim=0)
def _run_device_to_host_indexer_copy(self, io_backend: str):
page_size = 1 if is_hip() else 64
layer_num = 2
size = page_size * 4
device_pool = NSATokenToKVPool(
size=size,
page_size=page_size,
kv_lora_rank=128,
dtype=torch.bfloat16,
qk_rope_head_dim=32,
layer_num=layer_num,
device="cuda",
enable_memory_saver=False,
index_head_dim=128,
)
pin_memory = io_backend == "kernel"
original_alloc = ALLOC_MEMORY_FUNCS["cuda"]
if pin_memory:
ALLOC_MEMORY_FUNCS["cuda"] = alloc_with_pin_memory
try:
host_pool = NSATokenToKVPoolHost(
device_pool=device_pool,
host_to_device_ratio=2.0,
host_size=0,
page_size=page_size,
layout="layer_first",
pin_memory=pin_memory,
device="cpu",
)
finally:
ALLOC_MEMORY_FUNCS["cuda"] = original_alloc
for layer_id in range(layer_num):
buf = device_pool.index_k_with_scale_buffer[layer_id]
data = torch.arange(
buf.numel(), device=buf.device, dtype=torch.uint8
).view_as(buf)
buf.copy_((data + layer_id) % 256)
kv_buf = device_pool.kv_buffer[layer_id]
kv_data = torch.arange(
kv_buf.numel(), device=kv_buf.device, dtype=kv_buf.dtype
).view_as(kv_buf)
kv_buf.copy_(kv_data + layer_id)
device_pages = torch.tensor([1, 2, 3], device="cuda", dtype=torch.int64)
host_pages = torch.tensor(
[0, 1, 2],
device="cuda" if io_backend == "kernel" else "cpu",
dtype=torch.int64,
)
device_indices = self._token_indices_for_pages(
device_pages, page_size, device="cuda"
)
host_indices = self._token_indices_for_pages(
host_pages,
page_size,
device="cuda" if io_backend == "kernel" else "cpu",
)
host_pool.backup_from_device_all_layer(
device_pool, host_indices, device_indices, io_backend
)
for layer_id in range(layer_num):
for host_page, device_page in zip(
host_pages.tolist(), device_pages.tolist()
):
got = host_pool.index_k_with_scale_buffer[layer_id][host_page].cpu()
expected = device_pool.index_k_with_scale_buffer[layer_id][
device_page
].cpu()
self.assertTrue(torch.equal(got, expected))
host_start = host_page * page_size
device_start = device_page * page_size
got_kv = host_pool.kv_buffer[layer_id][
host_start : host_start + page_size
].cpu()
expected_kv = device_pool.kv_buffer[layer_id][
device_start : device_start + page_size
].cpu()
self.assertTrue(torch.equal(got_kv, expected_kv))
def test_device_to_host_indexer_kernel(self):
self._run_device_to_host_indexer_copy(io_backend="kernel")
def test_device_to_host_indexer_direct(self):
self._run_device_to_host_indexer_copy(io_backend="direct")
if __name__ == "__main__":
unittest.main()