From 180594358b990b2a5ce8140fb64aae90d73910fd Mon Sep 17 00:00:00 2001 From: zhangheng Date: Tue, 3 Feb 2026 10:07:37 +0800 Subject: [PATCH] [HiCache]: Support DeepSeek v32 cpu offloading (#17415) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: 晟海 --- python/sglang/srt/mem_cache/hiradix_cache.py | 16 +- .../sglang/srt/mem_cache/memory_pool_host.py | 220 +++++++++++++++--- .../hicache/test_nsa_pool_host_unit.py | 130 +++++++++++ 3 files changed, 334 insertions(+), 32 deletions(-) create mode 100644 test/registered/hicache/test_nsa_pool_host_unit.py diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index 7b016a144..793eca10a 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -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, diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index f40ddb17d..ba969cda5 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -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 + ) diff --git a/test/registered/hicache/test_nsa_pool_host_unit.py b/test/registered/hicache/test_nsa_pool_host_unit.py new file mode 100644 index 000000000..7c1f53f93 --- /dev/null +++ b/test/registered/hicache/test_nsa_pool_host_unit.py @@ -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()