From cb93e495334b5173a32f96915c96f3142d979755 Mon Sep 17 00:00:00 2001 From: hzh0425 Date: Thu, 14 May 2026 01:47:16 +0800 Subject: [PATCH 01/16] upd ci test --- .../test_unified_radix_hicache_kl.py | 53 +++++++++++++++++++ 1 file changed, 53 insertions(+) diff --git a/test/registered/radix_cache/test_unified_radix_hicache_kl.py b/test/registered/radix_cache/test_unified_radix_hicache_kl.py index 81253f6a7629..c6cb00e90237 100644 --- a/test/registered/radix_cache/test_unified_radix_hicache_kl.py +++ b/test/registered/radix_cache/test_unified_radix_hicache_kl.py @@ -33,6 +33,9 @@ DSV32_MODEL = "deepseek-ai/DeepSeek-V3.2" DSV32_LAUNCH_TIMEOUT = 3600 +GLM5_MODEL = "zai-org/GLM-5-FP8" +GLM5_LAUNCH_TIMEOUT = 3600 + register_cuda_ci(est_time=900, suite="nightly-8-gpu-h200", nightly=True) @@ -158,5 +161,55 @@ def tearDownClass(cls): kill_process_tree(cls.process.pid) +class TestUnifiedGLM5HiCache(UnifiedRadixTreeTestMixin, CustomTestCase): + """GLM-5 FP8 (DSA) + HiCache + UnifiedRadixCache.""" + + kl_threshold = 0.0035 + sampling_temperature = 0 + decode_cache_assert = staticmethod(_assert_dsv4_decode_cached_tokens) + gsm8k_threshold = 0.90 + num_gsm8k_questions = 100 + + @unittest.skip("no stable.") + def test_multiturn_logprobs_match(self): + pass + + @classmethod + def setUpClass(cls): + cls.model = GLM5_MODEL + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=GLM5_LAUNCH_TIMEOUT, + other_args=[ + "--trust-remote-code", + "--tp", + "8", + "--model-loader-extra-config", + '{"enable_multithread_load": true, "num_threads": 64}', + "--enable-hierarchical-cache", + "--hicache-ratio", + "4", + "--hicache-write-policy", + "write_through", + "--hicache-io-backend", + "direct", + "--hicache-mem-layout", + "page_first_direct", + "--max-total-tokens", + "20000", + "--max-running-requests", + "4", + ], + env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"}, + ) + cls.input_ids = get_input_ids(cls.model, num_samples=18) + + @classmethod + def tearDownClass(cls): + kill_process_tree(cls.process.pid) + + if __name__ == "__main__": unittest.main() From e93952cdf49aafd196857c32bbd11329cb68cb4a Mon Sep 17 00:00:00 2001 From: hzh0425 Date: Thu, 14 May 2026 02:01:26 +0800 Subject: [PATCH 02/16] upd test --- .../radix_cache/test_unified_radix_hicache_kl.py | 14 +++++--------- 1 file changed, 5 insertions(+), 9 deletions(-) diff --git a/test/registered/radix_cache/test_unified_radix_hicache_kl.py b/test/registered/radix_cache/test_unified_radix_hicache_kl.py index c6cb00e90237..ff3e49aaf63b 100644 --- a/test/registered/radix_cache/test_unified_radix_hicache_kl.py +++ b/test/registered/radix_cache/test_unified_radix_hicache_kl.py @@ -33,9 +33,6 @@ DSV32_MODEL = "deepseek-ai/DeepSeek-V3.2" DSV32_LAUNCH_TIMEOUT = 3600 -GLM5_MODEL = "zai-org/GLM-5-FP8" -GLM5_LAUNCH_TIMEOUT = 3600 - register_cuda_ci(est_time=900, suite="nightly-8-gpu-h200", nightly=True) @@ -161,13 +158,12 @@ def tearDownClass(cls): kill_process_tree(cls.process.pid) -class TestUnifiedGLM5HiCache(UnifiedRadixTreeTestMixin, CustomTestCase): - """GLM-5 FP8 (DSA) + HiCache + UnifiedRadixCache.""" +class TestUnifiedDeepSeekV32HiCache(UnifiedRadixTreeTestMixin, CustomTestCase): + """DeepSeek V3.2 (DSA) + HiCache + UnifiedRadixCache.""" kl_threshold = 0.0035 sampling_temperature = 0 - decode_cache_assert = staticmethod(_assert_dsv4_decode_cached_tokens) - gsm8k_threshold = 0.90 + gsm8k_threshold = 0.93 num_gsm8k_questions = 100 @unittest.skip("no stable.") @@ -176,12 +172,12 @@ def test_multiturn_logprobs_match(self): @classmethod def setUpClass(cls): - cls.model = GLM5_MODEL + cls.model = DSV32_MODEL cls.base_url = DEFAULT_URL_FOR_TEST cls.process = popen_launch_server( cls.model, cls.base_url, - timeout=GLM5_LAUNCH_TIMEOUT, + timeout=DSV32_LAUNCH_TIMEOUT, other_args=[ "--trust-remote-code", "--tp", From f67e742d835aa051385ddabe5eeae47c35f20ced Mon Sep 17 00:00:00 2001 From: hzh0425 Date: Thu, 14 May 2026 14:09:45 +0800 Subject: [PATCH 03/16] upd test --- .../test_unified_radix_hicache_kl.py | 49 ------------------- 1 file changed, 49 deletions(-) diff --git a/test/registered/radix_cache/test_unified_radix_hicache_kl.py b/test/registered/radix_cache/test_unified_radix_hicache_kl.py index ff3e49aaf63b..81253f6a7629 100644 --- a/test/registered/radix_cache/test_unified_radix_hicache_kl.py +++ b/test/registered/radix_cache/test_unified_radix_hicache_kl.py @@ -158,54 +158,5 @@ def tearDownClass(cls): kill_process_tree(cls.process.pid) -class TestUnifiedDeepSeekV32HiCache(UnifiedRadixTreeTestMixin, CustomTestCase): - """DeepSeek V3.2 (DSA) + HiCache + UnifiedRadixCache.""" - - kl_threshold = 0.0035 - sampling_temperature = 0 - gsm8k_threshold = 0.93 - num_gsm8k_questions = 100 - - @unittest.skip("no stable.") - def test_multiturn_logprobs_match(self): - pass - - @classmethod - def setUpClass(cls): - cls.model = DSV32_MODEL - cls.base_url = DEFAULT_URL_FOR_TEST - cls.process = popen_launch_server( - cls.model, - cls.base_url, - timeout=DSV32_LAUNCH_TIMEOUT, - other_args=[ - "--trust-remote-code", - "--tp", - "8", - "--model-loader-extra-config", - '{"enable_multithread_load": true, "num_threads": 64}', - "--enable-hierarchical-cache", - "--hicache-ratio", - "4", - "--hicache-write-policy", - "write_through", - "--hicache-io-backend", - "direct", - "--hicache-mem-layout", - "page_first_direct", - "--max-total-tokens", - "20000", - "--max-running-requests", - "4", - ], - env={"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1"}, - ) - cls.input_ids = get_input_ids(cls.model, num_samples=18) - - @classmethod - def tearDownClass(cls): - kill_process_tree(cls.process.pid) - - if __name__ == "__main__": unittest.main() From 5bd257b9faa63f3e723514f09c1c0b4c7b6c7c7a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=99=9F=E6=B5=B7?= Date: Thu, 14 May 2026 21:14:31 +0800 Subject: [PATCH 04/16] support deepseek v4 host pool layout Co-authored-by: hzh0425 --- .../hybrid_cache/hybrid_pool_assembler.py | 7 + .../sglang/srt/mem_cache/memory_pool_host.py | 441 +++++++++++++----- .../test_unified_radix_hicache_kl.py | 39 +- 3 files changed, 380 insertions(+), 107 deletions(-) diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index 3de18edde6d7..091afaad43cb 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -322,6 +322,7 @@ def build_deepseek_v4_hicache_stack( item_bytes=kvcache.swa_kv_pool.bytes_per_page_padded, num_host_pages=swa_num_host_pages, slot_page_size=kvcache.swa_page_size, + layout=server_args.hicache_mem_layout, allocator_type=server_args.hicache_storage_backend, ) swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator @@ -354,6 +355,7 @@ def build_deepseek_v4_hicache_stack( item_bytes=kvcache.c4_kv_pool.bytes_per_page_padded, num_host_pages=num_host_pages, slot_page_size=page_size, + layout=server_args.hicache_mem_layout, allocator_type=server_args.hicache_storage_backend, ) c4_indexer_host_pool = DeepSeekV4PagedHostPool( @@ -365,6 +367,7 @@ def build_deepseek_v4_hicache_stack( ), num_host_pages=num_host_pages, slot_page_size=page_size, + layout=server_args.hicache_mem_layout, allocator_type=server_args.hicache_storage_backend, ) c4_state_host_pool = DeepSeekV4StateHostPool( @@ -375,6 +378,7 @@ def build_deepseek_v4_hicache_stack( ], num_host_pages=swa_num_host_pages, swa_page_size=kvcache.swa_page_size, + layout=server_args.hicache_mem_layout, allocator_type=server_args.hicache_storage_backend, ) c4_indexer_state_host_pool = DeepSeekV4StateHostPool( @@ -385,6 +389,7 @@ def build_deepseek_v4_hicache_stack( ], num_host_pages=swa_num_host_pages, swa_page_size=kvcache.swa_page_size, + layout=server_args.hicache_mem_layout, allocator_type=server_args.hicache_storage_backend, ) entries.extend( @@ -427,6 +432,7 @@ def build_deepseek_v4_hicache_stack( item_bytes=kvcache.c128_kv_pool.bytes_per_page_padded, num_host_pages=num_host_pages, slot_page_size=page_size, + layout=server_args.hicache_mem_layout, allocator_type=server_args.hicache_storage_backend, ) c128_state_host_pool = DeepSeekV4StateHostPool( @@ -437,6 +443,7 @@ def build_deepseek_v4_hicache_stack( ], num_host_pages=swa_num_host_pages, swa_page_size=kvcache.swa_page_size, + layout=server_args.hicache_mem_layout, allocator_type=server_args.hicache_storage_backend, ) entries.extend( diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index c8d3b9a7bf55..cd2d1983765c 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -1754,6 +1754,7 @@ def __init__( item_bytes: int, num_host_pages: int, slot_page_size: int, + layout: str = "layer_first", device: str = "cpu", pin_memory: bool = True, allocator_type: str = "default", @@ -1769,7 +1770,7 @@ def __init__( self.allocator = get_allocator_from_storage(allocator_type) self.page_size = slot_page_size self.size = num_host_pages * slot_page_size - self.layout = "layer_first" + self.layout = layout self.size_per_token = item_bytes self.start_layer = 0 self.end_layer = self.layer_num @@ -1789,42 +1790,68 @@ def __init__( ) alloc_func = ALLOC_MEMORY_FUNCS[self.gpu_device] - self.kv_buffer = [ - alloc_func( - (num_host_pages, self.item_bytes), + if self.layout == "layer_first": + self.kv_buffer = [ + alloc_func( + (num_host_pages, self.item_bytes), + dtype=self.dtype, + device=self.device, + pin_memory=self.pin_memory, + allocator=self.allocator, + ) + for _ in range(self.layer_num) + ] + self.data_refs = [self.kv_buffer[i] for i in range(self.layer_num)] + elif self.layout == "page_first": + self.kv_buffer = alloc_func( + (num_host_pages, self.layer_num, self.item_bytes), dtype=self.dtype, device=self.device, pin_memory=self.pin_memory, allocator=self.allocator, ) - for _ in range(self.layer_num) - ] - self.data_refs = [self.kv_buffer[i] for i in range(self.layer_num)] + self.data_refs = [self.kv_buffer[:, i, :] for i in range(self.layer_num)] + elif self.layout == "page_first_direct": + self.kv_buffer = alloc_func( + (num_host_pages, self.layer_num, 1, self.item_bytes), + dtype=self.dtype, + device=self.device, + pin_memory=self.pin_memory, + allocator=self.allocator, + ) + self.data_refs = [] + else: + raise ValueError(f"Unsupported layout: {self.layout}") logger.info( "Allocating %.2f GB host memory for V4 paged pool '%s' " - "(layers=%d, pages=%d, item_bytes=%d).", + "(layers=%d, pages=%d, item_bytes=%d, layout=%s).", requested_bytes / 1e9, self.pool_name, self.layer_num, num_host_pages, self.item_bytes, + self.layout, ) - self.clear() - def _to_page_indices(self, indices: torch.Tensor) -> torch.Tensor: - if indices.numel() % self.slot_page_size != 0: - raise ValueError( - f"{self.pool_name} transfer indices must be page-aligned, " - f"got numel={indices.numel()}, slot_page_size={self.slot_page_size}" + self.device_ptrs = torch.tensor( + [x.data_ptr() for x in self.device_buffers], + dtype=torch.uint64, + device=self.gpu_device, + ) + self.data_ptrs = ( + torch.tensor( + [x.data_ptr() for x in self.data_refs], + dtype=torch.uint64, + device=self.gpu_device, ) - return indices.reshape(-1, self.slot_page_size)[:, 0] // self.slot_page_size + if self.data_refs + else None + ) + self.clear() - def _check_io_backend(self, io_backend: str) -> None: - if io_backend != "direct": - raise NotImplementedError( - f"{self.pool_name} supports only direct io_backend, got {io_backend}" - ) + def _to_page_indices(self, indices: torch.Tensor) -> torch.Tensor: + return torch.unique_consecutive(indices.to(torch.int64) // self.slot_page_size) def get_size_per_token(self): return self.item_bytes @@ -1836,7 +1863,7 @@ def init_kv_buffer(self): return self.kv_buffer def get_hybrid_pool_buffer(self): - return self.kv_buffer + return self.kv_buffer if isinstance(self.kv_buffer, list) else [self.kv_buffer] def clear(self): self.free_slots = torch.arange(self.size, dtype=torch.int64) @@ -1867,38 +1894,106 @@ def backup_from_device_all_layer( ): if host_indices is None or device_indices is None: return - self._check_io_backend(io_backend) host_rows = self._to_page_indices(host_indices) device_rows = self._to_page_indices(device_indices) - transfer_kv_direct( - src_layers=self.device_buffers, - dst_layers=self.data_refs, - src_indices=device_rows, - dst_indices=host_rows, - page_size=1, - ) + if io_backend == "kernel" and self.layout == "layer_first": + assert self.data_ptrs is not None + transfer_kv_all_layer_mla( + src_layers=self.device_ptrs, + dst_layers=self.data_ptrs, + src_indices=device_rows, + dst_indices=host_rows, + item_size=self.item_bytes, + num_layers=self.layer_num, + ) + elif io_backend == "kernel" and self.layout == "page_first": + transfer_kv_all_layer_mla_lf_pf( + src_layers=self.device_ptrs, + dst=self.kv_buffer, + src_indices=device_rows, + dst_indices=host_rows, + item_size=self.item_bytes, + dst_layout_dim=self.layer_num * self.item_bytes, + num_layers=self.layer_num, + ) + elif io_backend == "direct" and self.layout == "layer_first": + transfer_kv_direct( + src_layers=self.device_buffers, + dst_layers=self.data_refs, + src_indices=device_rows, + dst_indices=host_rows, + page_size=1, + ) + elif io_backend == "direct" and self.layout == "page_first_direct": + transfer_kv_all_layer_direct_lf_pf( + src_ptrs=self.device_buffers, + dst_ptrs=[self.kv_buffer], + src_indices=device_rows, + dst_indices=host_rows, + page_size=1, + ) + else: + raise ValueError( + f"Unsupported V4 paged host layout/backend: {self.layout}/{io_backend}" + ) def load_to_device_per_layer( self, device_pool, host_indices, device_indices, layer_id, io_backend ): if host_indices is None or device_indices is None: return - self._check_io_backend(io_backend) host_rows = self._to_page_indices(host_indices) device_rows = self._to_page_indices(device_indices) - transfer_kv_direct( - src_layers=[self.kv_buffer[layer_id]], - dst_layers=[self.device_buffers[layer_id]], - src_indices=host_rows, - dst_indices=device_rows, - page_size=1, - ) + if io_backend == "kernel" and self.layout == "layer_first": + transfer_kv_per_layer_mla( + src=self.data_refs[layer_id], + dst=self.device_buffers[layer_id], + src_indices=host_rows, + dst_indices=device_rows, + item_size=self.item_bytes, + ) + elif io_backend == "kernel" and self.layout == "page_first": + transfer_kv_per_layer_mla_pf_lf( + src=self.kv_buffer, + dst=self.device_buffers[layer_id], + src_indices=host_rows, + dst_indices=device_rows, + layer_id=layer_id, + item_size=self.item_bytes, + src_layout_dim=self.layer_num * self.item_bytes, + ) + elif io_backend == "direct" and self.layout == "layer_first": + transfer_kv_direct( + src_layers=[self.data_refs[layer_id]], + dst_layers=[self.device_buffers[layer_id]], + src_indices=host_rows, + dst_indices=device_rows, + page_size=1, + ) + elif io_backend == "direct" and self.layout == "page_first_direct": + transfer_kv_per_layer_direct_pf_lf( + src_ptrs=[self.kv_buffer], + dst_ptrs=[self.device_buffers[layer_id]], + src_indices=host_rows, + dst_indices=device_rows, + layer_id=layer_id, + page_size=1, + ) + else: + raise ValueError( + f"Unsupported V4 paged host layout/backend: {self.layout}/{io_backend}" + ) def get_data_page(self, index, flat=True): index = int(index) // self.slot_page_size - data_page = torch.stack( - [self.kv_buffer[i][index] for i in range(self.layer_num)] - ) + if self.layout == "layer_first": + data_page = torch.stack( + [self.kv_buffer[i][index] for i in range(self.layer_num)] + ) + elif self.layout in ["page_first", "page_first_direct"]: + data_page = self.kv_buffer[index] + else: + raise ValueError(f"Unsupported layout: {self.layout}") return data_page.flatten() if flat else data_page def get_dummy_flat_data_page(self): @@ -1911,22 +2006,41 @@ def get_dummy_flat_data_page(self): def set_from_flat_data_page(self, index, data_page): index = int(index) // self.slot_page_size - data = data_page.view(self.dtype).reshape(self.layer_num, self.item_bytes) - for i in range(self.layer_num): - self.kv_buffer[i][index].copy_(data[i]) + if self.layout == "layer_first": + data = data_page.view(self.dtype).reshape(self.layer_num, self.item_bytes) + for i in range(self.layer_num): + self.kv_buffer[i][index].copy_(data[i]) + elif self.layout == "page_first": + self.kv_buffer[index].copy_( + data_page.view(self.dtype).reshape(self.layer_num, self.item_bytes) + ) + elif self.layout == "page_first_direct": + self.kv_buffer[index].copy_( + data_page.view(self.dtype).reshape(self.layer_num, 1, self.item_bytes) + ) + else: + raise ValueError(f"Unsupported layout: {self.layout}") def get_page_buffer_meta(self, indices): ptr_list = [] rows = self._to_page_indices(indices).tolist() - for row in rows: - for layer_id in range(self.layer_num): - ptr = ( - self.kv_buffer[layer_id].data_ptr() - + int(row) * self.item_bytes * self.dtype.itemsize - ) - ptr_list.append(ptr) - element_size = self.item_bytes * self.dtype.itemsize - return ptr_list, [element_size] * len(ptr_list) + if self.layout == "layer_first": + for row in rows: + page_index = int(row) + for layer_id in range(self.layer_num): + ptr = ( + self.kv_buffer[layer_id].data_ptr() + + page_index * self.item_bytes * self.dtype.itemsize + ) + ptr_list.append(ptr) + element_size = self.item_bytes * self.dtype.itemsize + return ptr_list, [element_size] * len(ptr_list) + if self.layout in ["page_first", "page_first_direct"]: + page_bytes = self.layer_num * self.item_bytes * self.dtype.itemsize + for row in rows: + ptr_list.append(self.kv_buffer[int(row)].data_ptr()) + return ptr_list, [page_bytes] * len(ptr_list) + raise ValueError(f"Unsupported layout: {self.layout}") class DeepSeekV4StateHostPool(HostKVCache): @@ -1938,6 +2052,7 @@ def __init__( state_pools: list, num_host_pages: int, swa_page_size: int, + layout: str = "layer_first", device: str = "cpu", pin_memory: bool = True, allocator_type: str = "default", @@ -1956,7 +2071,7 @@ def __init__( self.allocator = get_allocator_from_storage(allocator_type) self.page_size = swa_page_size self.size = num_host_pages * swa_page_size - self.layout = "layer_first" + self.layout = layout self.start_layer = 0 self.end_layer = self.layer_num self.lock = threading.RLock() @@ -1979,25 +2094,61 @@ def __init__( ) alloc_func = ALLOC_MEMORY_FUNCS[self.gpu_device] - self.kv_buffer = [ - alloc_func( - (num_host_pages, self.state_page_bytes), + if self.layout == "layer_first": + self.kv_buffer = [ + alloc_func( + (num_host_pages, self.state_page_bytes), + dtype=self.dtype, + device=self.device, + pin_memory=self.pin_memory, + allocator=self.allocator, + ) + for _ in range(self.layer_num) + ] + self.data_refs = [self.kv_buffer[i] for i in range(self.layer_num)] + elif self.layout == "page_first": + self.kv_buffer = alloc_func( + (num_host_pages, self.layer_num, self.state_page_bytes), dtype=self.dtype, device=self.device, pin_memory=self.pin_memory, allocator=self.allocator, ) - for _ in range(self.layer_num) - ] - self.data_refs = [self.kv_buffer[i] for i in range(self.layer_num)] + self.data_refs = [self.kv_buffer[:, i, :] for i in range(self.layer_num)] + elif self.layout == "page_first_direct": + self.kv_buffer = alloc_func( + (num_host_pages, self.layer_num, 1, self.state_page_bytes), + dtype=self.dtype, + device=self.device, + pin_memory=self.pin_memory, + allocator=self.allocator, + ) + self.data_refs = [] + else: + raise ValueError(f"Unsupported layout: {self.layout}") logger.info( "Allocating %.2f GB host memory for V4 state pool '%s' " - "(layers=%d, pages=%d, state_page_bytes=%d).", + "(layers=%d, pages=%d, state_page_bytes=%d, layout=%s).", requested_bytes / 1e9, self.pool_name, self.layer_num, num_host_pages, self.state_page_bytes, + self.layout, + ) + self.device_ptrs = torch.tensor( + [x.data_ptr() for x in self.device_page_views], + dtype=torch.uint64, + device=self.gpu_device, + ) + self.data_ptrs = ( + torch.tensor( + [x.data_ptr() for x in self.data_refs], + dtype=torch.uint64, + device=self.gpu_device, + ) + if self.data_refs + else None ) def _init_device_page_views(self) -> None: @@ -2034,18 +2185,7 @@ def _init_device_page_views(self) -> None: self.state_page_bytes = expected_state_page_bytes or 0 def _to_page_indices(self, indices: torch.Tensor) -> torch.Tensor: - if indices.numel() % self.swa_page_size != 0: - raise ValueError( - f"{self.pool_name} transfer indices must be SWA-page-aligned, " - f"got numel={indices.numel()}, swa_page_size={self.swa_page_size}" - ) - return indices.reshape(-1, self.swa_page_size)[:, 0] // self.swa_page_size - - def _check_io_backend(self, io_backend: str) -> None: - if io_backend != "direct": - raise NotImplementedError( - f"{self.pool_name} supports only direct io_backend, got {io_backend}" - ) + return torch.unique_consecutive(indices.to(torch.int64) // self.swa_page_size) def get_size_per_token(self): return self.state_page_bytes @@ -2057,7 +2197,7 @@ def init_kv_buffer(self): return self.kv_buffer def get_hybrid_pool_buffer(self): - return self.kv_buffer + return self.kv_buffer if isinstance(self.kv_buffer, list) else [self.kv_buffer] def clear(self): pass @@ -2084,38 +2224,106 @@ def backup_from_device_all_layer( ): if host_indices is None or device_indices is None: return - self._check_io_backend(io_backend) host_rows = self._to_page_indices(host_indices) device_rows = self._to_page_indices(device_indices) - transfer_kv_direct( - src_layers=self.device_page_views, - dst_layers=self.data_refs, - src_indices=device_rows, - dst_indices=host_rows, - page_size=1, - ) + if io_backend == "kernel" and self.layout == "layer_first": + assert self.data_ptrs is not None + transfer_kv_all_layer_mla( + src_layers=self.device_ptrs, + dst_layers=self.data_ptrs, + src_indices=device_rows, + dst_indices=host_rows, + item_size=self.state_page_bytes, + num_layers=self.layer_num, + ) + elif io_backend == "kernel" and self.layout == "page_first": + transfer_kv_all_layer_mla_lf_pf( + src_layers=self.device_ptrs, + dst=self.kv_buffer, + src_indices=device_rows, + dst_indices=host_rows, + item_size=self.state_page_bytes, + dst_layout_dim=self.layer_num * self.state_page_bytes, + num_layers=self.layer_num, + ) + elif io_backend == "direct" and self.layout == "layer_first": + transfer_kv_direct( + src_layers=self.device_page_views, + dst_layers=self.data_refs, + src_indices=device_rows, + dst_indices=host_rows, + page_size=1, + ) + elif io_backend == "direct" and self.layout == "page_first_direct": + transfer_kv_all_layer_direct_lf_pf( + src_ptrs=self.device_page_views, + dst_ptrs=[self.kv_buffer], + src_indices=device_rows, + dst_indices=host_rows, + page_size=1, + ) + else: + raise ValueError( + f"Unsupported V4 state host layout/backend: {self.layout}/{io_backend}" + ) def load_to_device_per_layer( self, device_pool, host_indices, device_indices, layer_id, io_backend ): if host_indices is None or device_indices is None: return - self._check_io_backend(io_backend) host_rows = self._to_page_indices(host_indices) device_rows = self._to_page_indices(device_indices) - transfer_kv_direct( - src_layers=[self.kv_buffer[layer_id]], - dst_layers=[self.device_page_views[layer_id]], - src_indices=host_rows, - dst_indices=device_rows, - page_size=1, - ) + if io_backend == "kernel" and self.layout == "layer_first": + transfer_kv_per_layer_mla( + src=self.data_refs[layer_id], + dst=self.device_page_views[layer_id], + src_indices=host_rows, + dst_indices=device_rows, + item_size=self.state_page_bytes, + ) + elif io_backend == "kernel" and self.layout == "page_first": + transfer_kv_per_layer_mla_pf_lf( + src=self.kv_buffer, + dst=self.device_page_views[layer_id], + src_indices=host_rows, + dst_indices=device_rows, + layer_id=layer_id, + item_size=self.state_page_bytes, + src_layout_dim=self.layer_num * self.state_page_bytes, + ) + elif io_backend == "direct" and self.layout == "layer_first": + transfer_kv_direct( + src_layers=[self.data_refs[layer_id]], + dst_layers=[self.device_page_views[layer_id]], + src_indices=host_rows, + dst_indices=device_rows, + page_size=1, + ) + elif io_backend == "direct" and self.layout == "page_first_direct": + transfer_kv_per_layer_direct_pf_lf( + src_ptrs=[self.kv_buffer], + dst_ptrs=[self.device_page_views[layer_id]], + src_indices=host_rows, + dst_indices=device_rows, + layer_id=layer_id, + page_size=1, + ) + else: + raise ValueError( + f"Unsupported V4 state host layout/backend: {self.layout}/{io_backend}" + ) def get_data_page(self, index, flat=True): index = int(index) // self.swa_page_size - data_page = torch.stack( - [self.kv_buffer[i][index] for i in range(self.layer_num)] - ) + if self.layout == "layer_first": + data_page = torch.stack( + [self.kv_buffer[i][index] for i in range(self.layer_num)] + ) + elif self.layout in ["page_first", "page_first_direct"]: + data_page = self.kv_buffer[index] + else: + raise ValueError(f"Unsupported layout: {self.layout}") return data_page.flatten() if flat else data_page def get_dummy_flat_data_page(self): @@ -2128,22 +2336,47 @@ def get_dummy_flat_data_page(self): def set_from_flat_data_page(self, index, data_page): index = int(index) // self.swa_page_size - data = data_page.view(self.dtype).reshape(self.layer_num, self.state_page_bytes) - for i in range(self.layer_num): - self.kv_buffer[i][index].copy_(data[i]) + if self.layout == "layer_first": + data = data_page.view(self.dtype).reshape( + self.layer_num, self.state_page_bytes + ) + for i in range(self.layer_num): + self.kv_buffer[i][index].copy_(data[i]) + elif self.layout == "page_first": + self.kv_buffer[index].copy_( + data_page.view(self.dtype).reshape( + self.layer_num, self.state_page_bytes + ) + ) + elif self.layout == "page_first_direct": + self.kv_buffer[index].copy_( + data_page.view(self.dtype).reshape( + self.layer_num, 1, self.state_page_bytes + ) + ) + else: + raise ValueError(f"Unsupported layout: {self.layout}") def get_page_buffer_meta(self, indices): ptr_list = [] rows = self._to_page_indices(indices).tolist() - for row in rows: - for layer_id in range(self.layer_num): - ptr = ( - self.kv_buffer[layer_id].data_ptr() - + int(row) * self.state_page_bytes * self.dtype.itemsize - ) - ptr_list.append(ptr) - element_size = self.state_page_bytes * self.dtype.itemsize - return ptr_list, [element_size] * len(ptr_list) + if self.layout == "layer_first": + for row in rows: + page_index = int(row) + for layer_id in range(self.layer_num): + ptr = ( + self.kv_buffer[layer_id].data_ptr() + + page_index * self.state_page_bytes * self.dtype.itemsize + ) + ptr_list.append(ptr) + element_size = self.state_page_bytes * self.dtype.itemsize + return ptr_list, [element_size] * len(ptr_list) + if self.layout in ["page_first", "page_first_direct"]: + page_bytes = self.layer_num * self.state_page_bytes * self.dtype.itemsize + for row in rows: + ptr_list.append(self.kv_buffer[int(row)].data_ptr()) + return ptr_list, [page_bytes] * len(ptr_list) + raise ValueError(f"Unsupported layout: {self.layout}") @dataclass diff --git a/test/registered/radix_cache/test_unified_radix_hicache_kl.py b/test/registered/radix_cache/test_unified_radix_hicache_kl.py index 81253f6a7629..dc12d9909776 100644 --- a/test/registered/radix_cache/test_unified_radix_hicache_kl.py +++ b/test/registered/radix_cache/test_unified_radix_hicache_kl.py @@ -5,6 +5,7 @@ via KL divergence. """ +import os import unittest from test_unified_radix_cache_kl import UnifiedRadixTreeTestMixin @@ -27,7 +28,10 @@ MAMBA_CHUNK_SIZE = 64 MAMBA_TRACK_INTERVAL = 128 -DSV4_FLASH_MODEL = "sgl-project/DeepSeek-V4-Flash-FP8" +DSV4_FLASH_MODEL = os.getenv( + "SGLANG_TEST_DSV4_FLASH_MODEL", + "sgl-project/DeepSeek-V4-Flash-FP8", +) DSV4_FLASH_LAUNCH_TIMEOUT = 3600 DSV32_MODEL = "deepseek-ai/DeepSeek-V3.2" @@ -99,6 +103,8 @@ def _assert_dsv4_decode_cached_tokens(result, history_len, output_len, label): class TestUnifiedDeepSeekV4FlashHiCache(UnifiedRadixTreeTestMixin, CustomTestCase): """DeepSeek V4 Flash FP8 + HiCache + UnifiedRadixCache.""" + hicache_io_backend = "direct" + hicache_mem_layout = "layer_first" kl_threshold = 0.0035 sampling_temperature = 0 decode_cache_assert = staticmethod(_assert_dsv4_decode_cached_tokens) @@ -136,9 +142,9 @@ def setUpClass(cls): "--hicache-write-policy", "write_through", "--hicache-io-backend", - "direct", + cls.hicache_io_backend, "--hicache-mem-layout", - "page_first_direct", + cls.hicache_mem_layout, "--swa-full-tokens-ratio", "0.25", "--max-total-tokens", @@ -158,5 +164,32 @@ def tearDownClass(cls): kill_process_tree(cls.process.pid) +class TestUnifiedDeepSeekV4FlashHiCachePageFirstDirect( + TestUnifiedDeepSeekV4FlashHiCache +): + """DeepSeek V4 Flash HiCache layout smoke: page_first_direct + direct.""" + + hicache_io_backend = "direct" + hicache_mem_layout = "page_first_direct" + + +class TestUnifiedDeepSeekV4FlashHiCacheLayerFirstKernel( + TestUnifiedDeepSeekV4FlashHiCache +): + """DeepSeek V4 Flash HiCache layout smoke: layer_first + kernel.""" + + hicache_io_backend = "kernel" + hicache_mem_layout = "layer_first" + + +class TestUnifiedDeepSeekV4FlashHiCachePageFirstKernel( + TestUnifiedDeepSeekV4FlashHiCache +): + """DeepSeek V4 Flash HiCache layout smoke: page_first + kernel.""" + + hicache_io_backend = "kernel" + hicache_mem_layout = "page_first" + + if __name__ == "__main__": unittest.main() From 27944ca7549f913a1cf80ae3d6da4bef18e8471d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=99=9F=E6=B5=B7?= Date: Fri, 15 May 2026 11:39:11 +0800 Subject: [PATCH 05/16] opt --- .../registered/radix_cache/test_unified_radix_hicache_kl.py | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/test/registered/radix_cache/test_unified_radix_hicache_kl.py b/test/registered/radix_cache/test_unified_radix_hicache_kl.py index dc12d9909776..0af645cccdf8 100644 --- a/test/registered/radix_cache/test_unified_radix_hicache_kl.py +++ b/test/registered/radix_cache/test_unified_radix_hicache_kl.py @@ -5,7 +5,6 @@ via KL divergence. """ -import os import unittest from test_unified_radix_cache_kl import UnifiedRadixTreeTestMixin @@ -28,10 +27,7 @@ MAMBA_CHUNK_SIZE = 64 MAMBA_TRACK_INTERVAL = 128 -DSV4_FLASH_MODEL = os.getenv( - "SGLANG_TEST_DSV4_FLASH_MODEL", - "sgl-project/DeepSeek-V4-Flash-FP8", -) +DSV4_FLASH_MODEL = "sgl-project/DeepSeek-V4-Flash-FP8" DSV4_FLASH_LAUNCH_TIMEOUT = 3600 DSV32_MODEL = "deepseek-ai/DeepSeek-V3.2" From 228532a39cb15d74ee55e1500a5384bfdf0be5f0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=99=9F=E6=B5=B7?= Date: Fri, 15 May 2026 11:48:00 +0800 Subject: [PATCH 06/16] opt --- python/sglang/srt/mem_cache/memory_pool_host.py | 14 ++++++++++++-- 1 file changed, 12 insertions(+), 2 deletions(-) diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index cd2d1983765c..01e98687423f 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -1851,7 +1851,12 @@ def __init__( self.clear() def _to_page_indices(self, indices: torch.Tensor) -> torch.Tensor: - return torch.unique_consecutive(indices.to(torch.int64) // self.slot_page_size) + if indices.numel() % self.slot_page_size != 0: + raise ValueError( + f"{self.pool_name} transfer indices must be page-aligned, " + f"got numel={indices.numel()}, slot_page_size={self.slot_page_size}" + ) + return indices.reshape(-1, self.slot_page_size)[:, 0] // self.slot_page_size def get_size_per_token(self): return self.item_bytes @@ -2185,7 +2190,12 @@ def _init_device_page_views(self) -> None: self.state_page_bytes = expected_state_page_bytes or 0 def _to_page_indices(self, indices: torch.Tensor) -> torch.Tensor: - return torch.unique_consecutive(indices.to(torch.int64) // self.swa_page_size) + if indices.numel() % self.swa_page_size != 0: + raise ValueError( + f"{self.pool_name} transfer indices must be SWA-page-aligned, " + f"got numel={indices.numel()}, swa_page_size={self.swa_page_size}" + ) + return indices.reshape(-1, self.swa_page_size)[:, 0] // self.swa_page_size def get_size_per_token(self): return self.state_page_bytes From 673716661fe66280cf210d6491710368e4d7048b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=99=9F=E6=B5=B7?= Date: Fri, 15 May 2026 12:26:18 +0800 Subject: [PATCH 07/16] opt --- .../test_unified_radix_hicache_kl.py | 20 +------------------ 1 file changed, 1 insertion(+), 19 deletions(-) diff --git a/test/registered/radix_cache/test_unified_radix_hicache_kl.py b/test/registered/radix_cache/test_unified_radix_hicache_kl.py index 0af645cccdf8..fb6c77e0ca54 100644 --- a/test/registered/radix_cache/test_unified_radix_hicache_kl.py +++ b/test/registered/radix_cache/test_unified_radix_hicache_kl.py @@ -99,7 +99,7 @@ def _assert_dsv4_decode_cached_tokens(result, history_len, output_len, label): class TestUnifiedDeepSeekV4FlashHiCache(UnifiedRadixTreeTestMixin, CustomTestCase): """DeepSeek V4 Flash FP8 + HiCache + UnifiedRadixCache.""" - hicache_io_backend = "direct" + hicache_io_backend = "kernel" hicache_mem_layout = "layer_first" kl_threshold = 0.0035 sampling_temperature = 0 @@ -169,23 +169,5 @@ class TestUnifiedDeepSeekV4FlashHiCachePageFirstDirect( hicache_mem_layout = "page_first_direct" -class TestUnifiedDeepSeekV4FlashHiCacheLayerFirstKernel( - TestUnifiedDeepSeekV4FlashHiCache -): - """DeepSeek V4 Flash HiCache layout smoke: layer_first + kernel.""" - - hicache_io_backend = "kernel" - hicache_mem_layout = "layer_first" - - -class TestUnifiedDeepSeekV4FlashHiCachePageFirstKernel( - TestUnifiedDeepSeekV4FlashHiCache -): - """DeepSeek V4 Flash HiCache layout smoke: page_first + kernel.""" - - hicache_io_backend = "kernel" - hicache_mem_layout = "page_first" - - if __name__ == "__main__": unittest.main() From 20b7321ef974b83c5892ff36e595ace9150f4ffe Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=99=9F=E6=B5=B7?= Date: Fri, 15 May 2026 14:30:44 +0800 Subject: [PATCH 08/16] opt --- python/sglang/srt/mem_cache/memory_pool_host.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index 01e98687423f..b4624ab54a9f 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -1790,6 +1790,7 @@ def __init__( ) alloc_func = ALLOC_MEMORY_FUNCS[self.gpu_device] + self.data_refs = [] if self.layout == "layer_first": self.kv_buffer = [ alloc_func( @@ -1810,7 +1811,6 @@ def __init__( pin_memory=self.pin_memory, allocator=self.allocator, ) - self.data_refs = [self.kv_buffer[:, i, :] for i in range(self.layer_num)] elif self.layout == "page_first_direct": self.kv_buffer = alloc_func( (num_host_pages, self.layer_num, 1, self.item_bytes), @@ -1819,7 +1819,6 @@ def __init__( pin_memory=self.pin_memory, allocator=self.allocator, ) - self.data_refs = [] else: raise ValueError(f"Unsupported layout: {self.layout}") @@ -2099,6 +2098,7 @@ def __init__( ) alloc_func = ALLOC_MEMORY_FUNCS[self.gpu_device] + self.data_refs = [] if self.layout == "layer_first": self.kv_buffer = [ alloc_func( @@ -2119,7 +2119,6 @@ def __init__( pin_memory=self.pin_memory, allocator=self.allocator, ) - self.data_refs = [self.kv_buffer[:, i, :] for i in range(self.layer_num)] elif self.layout == "page_first_direct": self.kv_buffer = alloc_func( (num_host_pages, self.layer_num, 1, self.state_page_bytes), @@ -2128,7 +2127,6 @@ def __init__( pin_memory=self.pin_memory, allocator=self.allocator, ) - self.data_refs = [] else: raise ValueError(f"Unsupported layout: {self.layout}") logger.info( From 6d1eab194c9409e045ce372776de591d181a0121 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=99=9F=E6=B5=B7?= Date: Fri, 15 May 2026 19:41:27 +0800 Subject: [PATCH 09/16] fix ci --- .../radix_cache/test_unified_radix_hicache_kl.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/test/registered/radix_cache/test_unified_radix_hicache_kl.py b/test/registered/radix_cache/test_unified_radix_hicache_kl.py index fb6c77e0ca54..d139475aceeb 100644 --- a/test/registered/radix_cache/test_unified_radix_hicache_kl.py +++ b/test/registered/radix_cache/test_unified_radix_hicache_kl.py @@ -99,8 +99,9 @@ def _assert_dsv4_decode_cached_tokens(result, history_len, output_len, label): class TestUnifiedDeepSeekV4FlashHiCache(UnifiedRadixTreeTestMixin, CustomTestCase): """DeepSeek V4 Flash FP8 + HiCache + UnifiedRadixCache.""" - hicache_io_backend = "kernel" - hicache_mem_layout = "layer_first" + hicache_io_backend = "direct" + hicache_mem_layout = "page_first_direct" + max_running_requests = 4 kl_threshold = 0.0035 sampling_temperature = 0 decode_cache_assert = staticmethod(_assert_dsv4_decode_cached_tokens) @@ -146,7 +147,7 @@ def setUpClass(cls): "--max-total-tokens", "20000", "--max-running-requests", - "4", + str(cls.max_running_requests), ], env={ "SGLANG_DSV4_FP4_EXPERTS": "0", @@ -165,8 +166,9 @@ class TestUnifiedDeepSeekV4FlashHiCachePageFirstDirect( ): """DeepSeek V4 Flash HiCache layout smoke: page_first_direct + direct.""" - hicache_io_backend = "direct" - hicache_mem_layout = "page_first_direct" + hicache_io_backend = "kernel" + hicache_mem_layout = "layer_first" + max_running_requests = 4 if __name__ == "__main__": From 1df9b901fd8448b1bb08b1800301d27d99570dec Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=99=9F=E6=B5=B7?= Date: Fri, 15 May 2026 19:45:05 +0800 Subject: [PATCH 10/16] fix --- test/registered/radix_cache/test_unified_radix_hicache_kl.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/test/registered/radix_cache/test_unified_radix_hicache_kl.py b/test/registered/radix_cache/test_unified_radix_hicache_kl.py index 84837e8d978b..d1f063b6432e 100644 --- a/test/registered/radix_cache/test_unified_radix_hicache_kl.py +++ b/test/registered/radix_cache/test_unified_radix_hicache_kl.py @@ -300,7 +300,7 @@ class TestUnifiedDeepSeekV4FlashHiCachePageFirstDirect( hicache_io_backend = "kernel" hicache_mem_layout = "layer_first" - max_running_requests = 4 + max_running_requests = 2 if __name__ == "__main__": From d4af9a40e19d50d81d071f1530873bad7692fe1d Mon Sep 17 00:00:00 2001 From: hzh0425 Date: Fri, 15 May 2026 22:23:02 +0800 Subject: [PATCH 11/16] upd --- .../test_unified_radix_hicache_kl.py | 24 +++++++++---------- 1 file changed, 11 insertions(+), 13 deletions(-) diff --git a/test/registered/radix_cache/test_unified_radix_hicache_kl.py b/test/registered/radix_cache/test_unified_radix_hicache_kl.py index d1f063b6432e..6119ca4d3215 100644 --- a/test/registered/radix_cache/test_unified_radix_hicache_kl.py +++ b/test/registered/radix_cache/test_unified_radix_hicache_kl.py @@ -112,8 +112,7 @@ class TestUnifiedDeepSeekV4FlashHiCache(UnifiedRadixTreeTestMixin, CustomTestCas hicache_io_backend = "direct" hicache_mem_layout = "page_first_direct" - max_running_requests = 4 - kl_threshold = 0.0035 + kl_threshold = 0.005 sampling_temperature = 0 decode_cache_assert = staticmethod(_assert_dsv4_decode_cached_tokens) gsm8k_threshold = 0.90 @@ -158,7 +157,7 @@ def setUpClass(cls): "--max-total-tokens", "20000", "--max-running-requests", - str(cls.max_running_requests), + "2", ], env={ "SGLANG_DSV4_FP4_EXPERTS": "0", @@ -172,6 +171,15 @@ def tearDownClass(cls): kill_process_tree(cls.process.pid) +class TestUnifiedDeepSeekV4FlashHiCachePageFirstDirect( + TestUnifiedDeepSeekV4FlashHiCache +): + """DeepSeek V4 Flash HiCache layout smoke: page_first_direct + direct.""" + + hicache_io_backend = "kernel" + hicache_mem_layout = "layer_first" + + class GSM8KTwoPassMixin: """Mixin: run GSM8K twice with flush in between, verify accuracy diff. @@ -293,15 +301,5 @@ def tearDownClass(cls): shutil.rmtree(cls.hicache_dir, ignore_errors=True) -class TestUnifiedDeepSeekV4FlashHiCachePageFirstDirect( - TestUnifiedDeepSeekV4FlashHiCache -): - """DeepSeek V4 Flash HiCache layout smoke: page_first_direct + direct.""" - - hicache_io_backend = "kernel" - hicache_mem_layout = "layer_first" - max_running_requests = 2 - - if __name__ == "__main__": unittest.main() From 350bfedf67ac22e42b19dcddfafb048b5c29a308 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=99=9F=E6=B5=B7?= Date: Sat, 16 May 2026 23:39:49 +0800 Subject: [PATCH 12/16] fix ci --- python/sglang/test/kl_multiturn_utils.py | 68 ++++++++++++++----- .../test_unified_radix_cache_kl.py | 4 ++ .../test_unified_radix_hicache_kl.py | 5 +- 3 files changed, 58 insertions(+), 19 deletions(-) diff --git a/python/sglang/test/kl_multiturn_utils.py b/python/sglang/test/kl_multiturn_utils.py index bd21c502321b..9558ed464237 100644 --- a/python/sglang/test/kl_multiturn_utils.py +++ b/python/sglang/test/kl_multiturn_utils.py @@ -2,6 +2,7 @@ from __future__ import annotations +import time from typing import Callable from sglang.test.kl_test_utils import ( @@ -145,7 +146,13 @@ def _interleave_order(n: int, branches_per_group: int) -> list[int] | None: def _generate_maybe_interleaved( - base_url, inputs, max_new_tokens, order=None, sampling_temperature: float = 1 + base_url, + inputs, + max_new_tokens, + order=None, + sampling_temperature: float = 1, + request_batch_size: int | None = None, + inter_batch_delay_s: float = 0, ): """Generate with optional interleaved submission order. @@ -153,22 +160,42 @@ def _generate_maybe_interleaved( original order so the caller always sees results[i] corresponds to inputs[i]. """ + + def _generate_ordered(ordered_inputs): + if ( + request_batch_size is None + or request_batch_size <= 0 + or request_batch_size >= len(ordered_inputs) + ): + return _generate( + base_url, + ordered_inputs, + max_new_tokens, + return_logprob=True, + temperature=sampling_temperature, + ) + + results = [] + for start in range(0, len(ordered_inputs), request_batch_size): + end = start + request_batch_size + results.extend( + _generate( + base_url, + ordered_inputs[start:end], + max_new_tokens, + return_logprob=True, + temperature=sampling_temperature, + ) + ) + if end < len(ordered_inputs) and inter_batch_delay_s > 0: + time.sleep(inter_batch_delay_s) + return results + if order is None: - return _generate( - base_url, - inputs, - max_new_tokens, - return_logprob=True, - temperature=sampling_temperature, - ) + return _generate_ordered(inputs) + ordered = [inputs[i] for i in order] - results = _generate( - base_url, - ordered, - max_new_tokens, - return_logprob=True, - temperature=sampling_temperature, - ) + results = _generate_ordered(ordered) unordered = [None] * len(results) for idx, orig in enumerate(order): unordered[orig] = results[idx] @@ -423,6 +450,8 @@ def test_input_output_logprobs_match_decode_cache_hit_helper( branches_per_group: int = 0, replay_batch_size: int = 1, sampling_temperature: float = 1, + request_batch_size: int | None = None, + inter_batch_delay_s: float = 0, ): """Verify logprobs when decode cache is hit. @@ -453,12 +482,13 @@ def test_input_output_logprobs_match_decode_cache_hit_helper( # Turn 1: populate cache, no assertion, no interleaving _flush_cache(base_url) - results = _generate( + results = _generate_maybe_interleaved( base_url, first_turn_input_ids, max_new_tokens, - return_logprob=True, - temperature=sampling_temperature, + sampling_temperature=sampling_temperature, + request_batch_size=request_batch_size, + inter_batch_delay_s=inter_batch_delay_s, ) assert len(results) == n @@ -478,6 +508,8 @@ def test_input_output_logprobs_match_decode_cache_hit_helper( max_new_tokens, order, sampling_temperature=sampling_temperature, + request_batch_size=request_batch_size, + inter_batch_delay_s=inter_batch_delay_s, ) assert len(results) == n diff --git a/test/registered/radix_cache/test_unified_radix_cache_kl.py b/test/registered/radix_cache/test_unified_radix_cache_kl.py index e6d38c6a2634..c6dc6271ff05 100644 --- a/test/registered/radix_cache/test_unified_radix_cache_kl.py +++ b/test/registered/radix_cache/test_unified_radix_cache_kl.py @@ -49,6 +49,8 @@ class UnifiedRadixTreeTestMixin: prefill_cache_assert = None decode_cache_assert = None sampling_temperature: float = 1 + decode_hit_request_batch_size: int | None = None + decode_hit_inter_batch_delay_s: float = 0 gsm8k_threshold: float = 0.93 mmlu_threshold: float = 0.8 @@ -163,6 +165,8 @@ def test_multiturn_decode_cache_hit_branching(self): branches_per_group=branches, max_new_tokens=self.max_new_tokens, sampling_temperature=self.sampling_temperature, + request_batch_size=self.decode_hit_request_batch_size, + inter_batch_delay_s=self.decode_hit_inter_batch_delay_s, ) diff --git a/test/registered/radix_cache/test_unified_radix_hicache_kl.py b/test/registered/radix_cache/test_unified_radix_hicache_kl.py index 6119ca4d3215..52bc08f40707 100644 --- a/test/registered/radix_cache/test_unified_radix_hicache_kl.py +++ b/test/registered/radix_cache/test_unified_radix_hicache_kl.py @@ -112,6 +112,7 @@ class TestUnifiedDeepSeekV4FlashHiCache(UnifiedRadixTreeTestMixin, CustomTestCas hicache_io_backend = "direct" hicache_mem_layout = "page_first_direct" + max_running_requests = 4 kl_threshold = 0.005 sampling_temperature = 0 decode_cache_assert = staticmethod(_assert_dsv4_decode_cached_tokens) @@ -157,7 +158,7 @@ def setUpClass(cls): "--max-total-tokens", "20000", "--max-running-requests", - "2", + str(cls.max_running_requests), ], env={ "SGLANG_DSV4_FP4_EXPERTS": "0", @@ -178,6 +179,8 @@ class TestUnifiedDeepSeekV4FlashHiCachePageFirstDirect( hicache_io_backend = "kernel" hicache_mem_layout = "layer_first" + decode_hit_request_batch_size = 4 + decode_hit_inter_batch_delay_s = 2 class GSM8KTwoPassMixin: From d500d494ce610ca8ff2e6e75a152b5239d3c7158 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=99=9F=E6=B5=B7?= Date: Sat, 16 May 2026 23:45:48 +0800 Subject: [PATCH 13/16] fix ci --- test/registered/radix_cache/test_unified_radix_hicache_kl.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/test/registered/radix_cache/test_unified_radix_hicache_kl.py b/test/registered/radix_cache/test_unified_radix_hicache_kl.py index 52bc08f40707..14e8c87f8102 100644 --- a/test/registered/radix_cache/test_unified_radix_hicache_kl.py +++ b/test/registered/radix_cache/test_unified_radix_hicache_kl.py @@ -115,6 +115,8 @@ class TestUnifiedDeepSeekV4FlashHiCache(UnifiedRadixTreeTestMixin, CustomTestCas max_running_requests = 4 kl_threshold = 0.005 sampling_temperature = 0 + decode_hit_request_batch_size = 4 + decode_hit_inter_batch_delay_s = 0.5 decode_cache_assert = staticmethod(_assert_dsv4_decode_cached_tokens) gsm8k_threshold = 0.90 num_gsm8k_questions = 100 @@ -179,8 +181,6 @@ class TestUnifiedDeepSeekV4FlashHiCachePageFirstDirect( hicache_io_backend = "kernel" hicache_mem_layout = "layer_first" - decode_hit_request_batch_size = 4 - decode_hit_inter_batch_delay_s = 2 class GSM8KTwoPassMixin: From 8779f6d5c488da2670c2b0d53ede40586c741878 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=99=9F=E6=B5=B7?= Date: Sun, 17 May 2026 18:16:39 +0800 Subject: [PATCH 14/16] fix ci --- python/sglang/test/kl_multiturn_utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/python/sglang/test/kl_multiturn_utils.py b/python/sglang/test/kl_multiturn_utils.py index 9558ed464237..56dc0db9f455 100644 --- a/python/sglang/test/kl_multiturn_utils.py +++ b/python/sglang/test/kl_multiturn_utils.py @@ -187,7 +187,7 @@ def _generate_ordered(ordered_inputs): temperature=sampling_temperature, ) ) - if end < len(ordered_inputs) and inter_batch_delay_s > 0: + if inter_batch_delay_s > 0: time.sleep(inter_batch_delay_s) return results From 99de29f987a122e1bf5aab663a2ef3699558973c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=99=9F=E6=B5=B7?= Date: Mon, 18 May 2026 10:28:55 +0800 Subject: [PATCH 15/16] fix ci --- .../radix_cache/test_unified_radix_cache_kl_hicache.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/test/registered/radix_cache/test_unified_radix_cache_kl_hicache.py b/test/registered/radix_cache/test_unified_radix_cache_kl_hicache.py index b86acdf519ed..d43b6fffc66b 100644 --- a/test/registered/radix_cache/test_unified_radix_cache_kl_hicache.py +++ b/test/registered/radix_cache/test_unified_radix_cache_kl_hicache.py @@ -97,7 +97,7 @@ class TestUnifiedDeepSeekV4FlashHiCache(UnifiedRadixTreeTestMixin, CustomTestCas max_running_requests = 4 kl_threshold = 0.005 sampling_temperature = 0 - decode_hit_request_batch_size = 4 + decode_hit_request_batch_size = 3 decode_hit_inter_batch_delay_s = 0.5 decode_cache_assert = staticmethod(_assert_dsv4_decode_cached_tokens) gsm8k_threshold = 0.90 From 0de1a5385e4bd8623ea4329d655f7f424ff40421 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=99=9F=E6=B5=B7?= Date: Mon, 18 May 2026 14:13:17 +0800 Subject: [PATCH 16/16] opt --- python/sglang/test/kl_multiturn_utils.py | 47 +++++++++--------------- 1 file changed, 18 insertions(+), 29 deletions(-) diff --git a/python/sglang/test/kl_multiturn_utils.py b/python/sglang/test/kl_multiturn_utils.py index 56dc0db9f455..97219b6d3ef7 100644 --- a/python/sglang/test/kl_multiturn_utils.py +++ b/python/sglang/test/kl_multiturn_utils.py @@ -160,42 +160,31 @@ def _generate_maybe_interleaved( original order so the caller always sees results[i] corresponds to inputs[i]. """ - - def _generate_ordered(ordered_inputs): - if ( - request_batch_size is None - or request_batch_size <= 0 - or request_batch_size >= len(ordered_inputs) - ): - return _generate( + ordered = inputs if order is None else [inputs[i] for i in order] + if not ordered: + return [] + + batch_size = ( + request_batch_size + if request_batch_size is not None and request_batch_size > 0 + else len(ordered) + ) + results = [] + for start in range(0, len(ordered), batch_size): + results.extend( + _generate( base_url, - ordered_inputs, + ordered[start : start + batch_size], max_new_tokens, return_logprob=True, temperature=sampling_temperature, ) - - results = [] - for start in range(0, len(ordered_inputs), request_batch_size): - end = start + request_batch_size - results.extend( - _generate( - base_url, - ordered_inputs[start:end], - max_new_tokens, - return_logprob=True, - temperature=sampling_temperature, - ) - ) - if inter_batch_delay_s > 0: - time.sleep(inter_batch_delay_s) - return results + ) + if batch_size < len(ordered) and inter_batch_delay_s > 0: + time.sleep(inter_batch_delay_s) if order is None: - return _generate_ordered(inputs) - - ordered = [inputs[i] for i in order] - results = _generate_ordered(ordered) + return results unordered = [None] * len(results) for idx, orig in enumerate(order): unordered[orig] = results[idx]