diff --git a/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py b/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py index bccdd809b..6321efb2e 100644 --- a/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py +++ b/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py @@ -19,8 +19,8 @@ from sglang.srt.mem_cache.memory_pool import ( ReqToTokenPool, ) from sglang.srt.mem_cache.memory_pool_host import ( - MHATokenToKVPoolHost, MLATokenToKVPoolHost, + get_mha_host_pool_cls, ) from sglang.srt.server_args import ServerArgs from sglang.srt.utils.common import ceil_align @@ -57,7 +57,7 @@ class DecodeKVCacheOffloadManager: ) kv_cache = self.token_to_kv_pool_allocator.get_kvcache() if isinstance(kv_cache, MHATokenToKVPool): - self.decode_host_mem_pool = MHATokenToKVPoolHost( + self.decode_host_mem_pool = get_mha_host_pool_cls(kv_cache)( kv_cache, server_args.hicache_ratio, server_args.hicache_size, diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index 463731f00..7fb56f943 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -44,8 +44,8 @@ from sglang.srt.mem_cache.memory_pool import ( MLATokenToKVPool, ) from sglang.srt.mem_cache.memory_pool_host import ( - MHATokenToKVPoolHost, MLATokenToKVPoolHost, + get_mha_host_pool_cls, ) from sglang.srt.mem_cache.radix_cache import ( RadixCache, @@ -78,7 +78,7 @@ class HiRadixCache(RadixCache): self.kv_cache = params.token_to_kv_pool_allocator.get_kvcache() if isinstance(self.kv_cache, MHATokenToKVPool): - self.token_to_kv_pool_host = MHATokenToKVPoolHost( + self.token_to_kv_pool_host = get_mha_host_pool_cls(self.kv_cache)( self.kv_cache, server_args.hicache_ratio, server_args.hicache_size, 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 2aac64b13..8c94c53b8 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 @@ -19,9 +19,9 @@ from sglang.srt.mem_cache.memory_pool_host import ( HostPoolGroup, LogicalHostPool, MambaPoolHost, - MHATokenToKVPoolHost, MLATokenToKVPoolHost, PoolEntry, + get_mha_host_pool_cls, ) from sglang.srt.mem_cache.unified_cache_components import ComponentType @@ -57,7 +57,9 @@ def build_kv_host_pool( use_mla: bool, override_kv_cache_dim: Optional[int] = None, ): - kv_host_pool_cls = MLATokenToKVPoolHost if use_mla else MHATokenToKVPoolHost + kv_host_pool_cls = ( + MLATokenToKVPoolHost if use_mla else get_mha_host_pool_cls(kv_pool) + ) kwargs = {} if override_kv_cache_dim is not None: kwargs["override_kv_cache_dim"] = override_kv_cache_dim diff --git a/python/sglang/srt/mem_cache/kv_cache_builder.py b/python/sglang/srt/mem_cache/kv_cache_builder.py index 2d2c406e1..726c5da39 100644 --- a/python/sglang/srt/mem_cache/kv_cache_builder.py +++ b/python/sglang/srt/mem_cache/kv_cache_builder.py @@ -89,8 +89,8 @@ def maybe_register_hicache_draft( MLATokenToKVPool, ) from sglang.srt.mem_cache.memory_pool_host import ( - MHATokenToKVPoolHost, MLATokenToKVPoolHost, + get_mha_host_pool_cls, ) pool = draft_kv_pool @@ -107,7 +107,7 @@ def maybe_register_hicache_draft( layout=server_args.hicache_mem_layout, ) if isinstance(pool, MHATokenToKVPool): - draft_host_pool = MHATokenToKVPoolHost(pool, **kw) + draft_host_pool = get_mha_host_pool_cls(pool)(pool, **kw) elif isinstance(pool, MLATokenToKVPool): draft_host_pool = MLATokenToKVPoolHost(pool, **kw) else: diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index b42ad7ba6..1e6596011 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -177,8 +177,21 @@ def get_allocator_from_storage(allocator_type): return HostTensorAllocator() +def _cuda_host_register(buffer: torch.Tensor) -> None: + cudart = torch.cuda.cudart() + n_bytes = buffer.numel() * buffer.element_size() + rc = cudart.cudaHostRegister(buffer.data_ptr(), n_bytes, 0) + if int(rc) != 0: + raise RuntimeError( + f"cudaHostRegister failed (rc={int(rc)}, " + f"{cudart.cudaGetErrorString(rc)}) for ptr={buffer.data_ptr():#x} " + f"size={n_bytes}; host buffer is not pinned and device transfers " + f"may silently return stale data." + ) + + def alloc_with_host_register( - dims, + dims: tuple, dtype: torch.dtype, device: str, pin_memory: bool, @@ -190,21 +203,12 @@ def alloc_with_host_register( """ buffer = allocator.allocate(dims, dtype=dtype, device=device) if pin_memory: - cudart = torch.cuda.cudart() - n_bytes = buffer.numel() * buffer.element_size() - rc = cudart.cudaHostRegister(buffer.data_ptr(), n_bytes, 0) - if int(rc) != 0: - raise RuntimeError( - f"cudaHostRegister failed (rc={int(rc)}, " - f"{cudart.cudaGetErrorString(rc)}) for ptr={buffer.data_ptr():#x} " - f"size={n_bytes}; host buffer is not pinned and device transfers " - f"may silently return stale data." - ) + _cuda_host_register(buffer) return buffer def alloc_with_pin_memory( - dims, + dims: tuple, dtype: torch.dtype, device: str, pin_memory: bool, @@ -429,7 +433,6 @@ class MHATokenToKVPoolHost(HostKVCache): self.head_num = self.device_pool.head_num self.head_dim = self.device_pool.head_dim self.layer_num = self.device_pool.layer_num - return self.head_dim * self.head_num * self.layer_num * self.dtype.itemsize * 2 def get_ksize_per_token(self): @@ -810,7 +813,7 @@ class MHATokenToKVPoolHost(HostKVCache): return ptr_list, element_size_list def get_page_buffer_meta(self, indices): - """ " + """ meta data for zero copy """ assert len(indices) % self.page_size == 0 @@ -896,6 +899,259 @@ class MHATokenToKVPoolHost(HostKVCache): return base_aligned and stride % page_size_bytes == 0 +class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost): + """Host KV pool for MHA models whose K and V have different head dims + (``head_dim != v_head_dim``), e.g. MiMo-V2. + + K and V are stored in two independent host buffers (``self.k_buffer`` and + ``self.v_buffer``) instead of a single ``(2, ...)`` tensor, so each side + keeps its native stride. The kernel transfer path dispatches K and V as + independent single-buffer copies so each side uses its own ``item_size``. + Direct transfer and the flat-page L3 storage interface assume a single + shared ``item_size`` in paths that are not safe for asymmetric K/V, so they + raise instead of silently corrupting V copies. + """ + + def get_size_per_token(self): + self.head_num = self.device_pool.head_num + self.head_dim = self.device_pool.head_dim + self.layer_num = self.device_pool.layer_num + self.v_head_dim = self.device_pool.v_head_dim + return ( + (self.head_dim + self.v_head_dim) + * self.head_num + * self.layer_num + * self.dtype.itemsize + ) + + def get_ksize_per_token(self): + return self.head_dim * self.head_num * self.layer_num * self.dtype.itemsize + + def init_kv_buffer(self): + if self.layout == "page_first": + k_dims = (self.size, self.layer_num, self.head_num, self.head_dim) + v_dims = (self.size, self.layer_num, self.head_num, self.v_head_dim) + else: + raise ValueError( + f"Unsupported layout for models with head_dim != v_head_dim: " + f"{self.layout}; expected 'page_first'." + ) + + # token_stride_size / layout_dim are intentionally NOT set: K and V + # have different strides, so any caller that reaches for a single + # shared stride is a bug. Such callers will fail loudly with + # AttributeError rather than silently use the K stride for V copies. + + alloc_func = ALLOC_MEMORY_FUNCS[self.device_pool.device] + k_buffer = alloc_func( + k_dims, + dtype=self.dtype, + device=self.device, + pin_memory=self.pin_memory, + allocator=self.allocator, + ) + v_buffer = alloc_func( + v_dims, + dtype=self.dtype, + device=self.device, + pin_memory=self.pin_memory, + allocator=self.allocator, + ) + return (k_buffer, v_buffer) + + def _k_token_stride_size(self) -> int: + return self.head_num * self.head_dim * self.dtype.itemsize + + def _v_token_stride_size(self) -> int: + return self.head_num * self.v_head_dim * self.dtype.itemsize + + def _k_layout_dim(self) -> int: + return self._k_token_stride_size() * self.layer_num + + def _v_layout_dim(self) -> int: + return self._v_token_stride_size() * self.layer_num + + def _flat_page_unsupported(self) -> NotImplementedError: + return NotImplementedError( + "Models with head_dim != v_head_dim do not support the flat-page " + "interface used by HiCache L3 storage backends {hf3fs, eic, nixl}. " + "Use a backend that does not use this interface (e.g. mooncake, simm)." + ) + + def load_to_device_per_layer( + self, + device_pool, + host_indices, + device_indices, + layer_id, + io_backend, + ): + if io_backend == "kernel": + if self.layout != "page_first": + raise ValueError( + f"Unsupported layout for models with head_dim != v_head_dim " + f"and io_backend='kernel': {self.layout}; expected 'page_first'." + ) + transfer_kv_per_layer_mla_pf_lf( + src=self.k_buffer, + dst=device_pool.k_buffer[layer_id], + src_indices=host_indices, + dst_indices=device_indices, + layer_id=layer_id, + item_size=self._k_token_stride_size(), + src_layout_dim=self._k_layout_dim(), + ) + transfer_kv_per_layer_mla_pf_lf( + src=self.v_buffer, + dst=device_pool.v_buffer[layer_id], + src_indices=host_indices, + dst_indices=device_indices, + layer_id=layer_id, + item_size=self._v_token_stride_size(), + src_layout_dim=self._v_layout_dim(), + ) + else: + raise ValueError( + f"Unsupported IO backend for models with head_dim != v_head_dim: " + f"{io_backend}; expected 'kernel'." + ) + + def backup_from_device_all_layer( + self, device_pool, host_indices, device_indices, io_backend + ): + if io_backend == "kernel": + if self.layout != "page_first": + raise ValueError( + f"Unsupported layout for models with head_dim != v_head_dim " + f"and io_backend='kernel': {self.layout}; expected 'page_first'." + ) + transfer_kv_all_layer_mla_lf_pf( + src_layers=device_pool.k_data_ptrs, + dst=self.k_buffer, + src_indices=device_indices, + dst_indices=host_indices, + item_size=self._k_token_stride_size(), + dst_layout_dim=self._k_layout_dim(), + num_layers=self.layer_num, + ) + transfer_kv_all_layer_mla_lf_pf( + src_layers=device_pool.v_data_ptrs, + dst=self.v_buffer, + src_indices=device_indices, + dst_indices=host_indices, + item_size=self._v_token_stride_size(), + dst_layout_dim=self._v_layout_dim(), + num_layers=self.layer_num, + ) + else: + raise ValueError( + f"Unsupported IO backend for models with head_dim != v_head_dim: " + f"{io_backend}; expected 'kernel'." + ) + + def get_data_page(self, index, flat: bool = True) -> torch.Tensor: + raise self._flat_page_unsupported() + + def get_dummy_flat_data_page(self) -> torch.Tensor: + raise self._flat_page_unsupported() + + def set_from_flat_data_page(self, index: int, data_page: torch.Tensor) -> None: + raise self._flat_page_unsupported() + + def get_split_heads_page_buffer_meta( + self, indices: torch.Tensor, split_factor: int + ): + raise NotImplementedError( + "get_split_heads_page_buffer_meta requires layout='page_head', " + "which is not supported for models with head_dim != v_head_dim." + ) + + def get_page_buffer_meta(self, indices): + assert len(indices) % self.page_size == 0 + if self.layout != "page_first": + raise ValueError( + f"Unsupported layout for models with head_dim != v_head_dim: " + f"{self.layout}" + ) + indices = indices.tolist() + k_base_ptr = self.k_buffer.data_ptr() + v_base_ptr = self.v_buffer.data_ptr() + k_element_size = ( + self.layer_num + * self.dtype.itemsize + * self.page_size + * self.head_num + * self.head_dim + ) + v_element_size = ( + self.layer_num + * self.dtype.itemsize + * self.page_size + * self.head_num + * self.v_head_dim + ) + ptr_list = [] + element_size_list = [] + for index in range(0, len(indices), self.page_size): + k_ptr = ( + k_base_ptr + + indices[index] + * self.layer_num + * self.head_num + * self.head_dim + * self.dtype.itemsize + ) + v_ptr = ( + v_base_ptr + + indices[index] + * self.layer_num + * self.head_num + * self.v_head_dim + * self.dtype.itemsize + ) + ptr_list.extend([k_ptr, v_ptr]) + element_size_list.extend([k_element_size, v_element_size]) + return ptr_list, element_size_list + + def is_stride_page_aligned(self, page_size_bytes: int = 4096) -> bool: + if self.layout != "page_first": + return False + k_stride = ( + self.page_size + * self.layer_num + * self.head_num + * self.head_dim + * self.dtype.itemsize + ) + v_stride = ( + self.page_size + * self.layer_num + * self.head_num + * self.v_head_dim + * self.dtype.itemsize + ) + base_aligned = ( + self.k_buffer.data_ptr() % page_size_bytes == 0 + and self.v_buffer.data_ptr() % page_size_bytes == 0 + ) + return ( + base_aligned + and k_stride % page_size_bytes == 0 + and v_stride % page_size_bytes == 0 + ) + + +def get_mha_host_pool_cls(device_pool: MHATokenToKVPool) -> type: + """Pick the right MHA host-pool class based on the device pool's K/V dims. + + Returns ``AsymmetricMHATokenToKVPoolHost`` when ``head_dim != v_head_dim`` + (e.g. MiMo-V2), else the default ``MHATokenToKVPoolHost``. + """ + if device_pool.head_dim != device_pool.v_head_dim: + return AsymmetricMHATokenToKVPoolHost + return MHATokenToKVPoolHost + + class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): device_pool: MLATokenToKVPool @@ -1256,7 +1512,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): raise ValueError(f"Unsupported layout: {self.layout}") def get_page_buffer_meta(self, indices): - """ " + """ meta data for zero copy """ assert len(indices) % self.page_size == 0 diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index a15a62cb1..20d9e427d 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -2418,14 +2418,29 @@ class ServerArgs: ) if self.enable_hierarchical_cache: - self.swa_full_tokens_ratio = 1.0 - logger.warning( - "Reset swa_full_tokens_ratio to 1.0 for MiMoV2 model with hierarchical cache" - ) - self.disable_hybrid_swa_memory = True - logger.warning( - "Disable hybrid SWA memory for MiMoV2 model with hierarchical cache" - ) + if not envs.SGLANG_ENABLE_UNIFIED_RADIX_TREE.get(): + raise ValueError( + "Hierarchical cache for MiMoV2 requires the unified " + "radix tree. Set SGLANG_ENABLE_UNIFIED_RADIX_TREE=1 " + "to enable --enable-hierarchical-cache for this model." + ) + + # MiMoV2 has head_dim != v_head_dim, so the host KV pool uses + # asymmetric K/V allocation. Only the kernel/page_first transfer + # path has a safe split K/V implementation. + if self.hicache_io_backend != "kernel": + logger.warning( + f"Force hicache_io_backend to 'kernel' for MiMoV2 model " + f"(was {self.hicache_io_backend!r})." + ) + self.hicache_io_backend = "kernel" + if self.hicache_mem_layout != "page_first": + logger.warning( + f"Force hicache_mem_layout to 'page_first' for " + f"MiMoV2 model (was {self.hicache_mem_layout!r}); " + f"asymmetric K/V HiCache requires kernel/page_first." + ) + self.hicache_mem_layout = "page_first" elif ( "Step3p5ForCausalLM" in model_arch or "Step3p7ForConditionalGeneration" in model_arch diff --git a/test/registered/jit/test_kvcacheio_asymmetric.py b/test/registered/jit/test_kvcacheio_asymmetric.py new file mode 100644 index 000000000..6d06ea681 --- /dev/null +++ b/test/registered/jit/test_kvcacheio_asymmetric.py @@ -0,0 +1,153 @@ +import sys +from types import SimpleNamespace + +import pytest +import torch + +from sglang.srt.mem_cache.memory_pool_host import AsymmetricMHATokenToKVPoolHost +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=10, suite="base-b-kernel-unit-1-gpu-large") + +# These tests use AsymmetricMHATokenToKVPoolHost methods and let that class call +# the real sgl-kernel transfer ops. The asymmetric host pool is kernel-only; +# direct/page_first_direct is intentionally rejected in the CPU dispatch tests. +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available(), reason="asymmetric host-pool tests require CUDA." +) + +DEVICE = "cuda" +PAGE_SIZE = 16 +NUM_LAYERS = 3 +TOTAL_ITEMS = PAGE_SIZE * 8 +HEAD_NUM = 4 +K_HEAD_DIM = 192 +V_HEAD_DIM = 128 +DTYPES = [torch.float16, torch.bfloat16] + + +def token_indices_for_pages(pages, page_size=PAGE_SIZE, device=None): + indices = torch.cat( + [ + torch.arange( + int(page) * page_size, + (int(page) + 1) * page_size, + dtype=torch.int64, + ) + for page in pages.tolist() + ] + ) + return indices if device is None else indices.to(device) + + +def fill_with_offset(tensor, offset): + data = torch.arange(tensor.numel(), device=tensor.device, dtype=tensor.dtype) + tensor.copy_((data + offset).view_as(tensor)) + + +def make_host_pool(dtype): + host = AsymmetricMHATokenToKVPoolHost.__new__(AsymmetricMHATokenToKVPoolHost) + host.layout = "page_first" + host.page_size = PAGE_SIZE + host.layer_num = NUM_LAYERS + host.head_num = HEAD_NUM + host.head_dim = K_HEAD_DIM + host.v_head_dim = V_HEAD_DIM + host.dtype = dtype + host.kv_buffer = ( + torch.zeros( + TOTAL_ITEMS, NUM_LAYERS, HEAD_NUM, K_HEAD_DIM, dtype=dtype + ).pin_memory(), + torch.zeros( + TOTAL_ITEMS, NUM_LAYERS, HEAD_NUM, V_HEAD_DIM, dtype=dtype + ).pin_memory(), + ) + return host + + +def make_device_pool(dtype): + k_buffer = [ + torch.empty(TOTAL_ITEMS, HEAD_NUM, K_HEAD_DIM, dtype=dtype, device=DEVICE) + for _ in range(NUM_LAYERS) + ] + v_buffer = [ + torch.empty(TOTAL_ITEMS, HEAD_NUM, V_HEAD_DIM, dtype=dtype, device=DEVICE) + for _ in range(NUM_LAYERS) + ] + for layer_id in range(NUM_LAYERS): + fill_with_offset(k_buffer[layer_id], layer_id * 1000) + fill_with_offset(v_buffer[layer_id], layer_id * 1000 + 100) + + return SimpleNamespace( + k_buffer=k_buffer, + v_buffer=v_buffer, + k_data_ptrs=torch.tensor( + [x.data_ptr() for x in k_buffer], dtype=torch.uint64, device=DEVICE + ), + v_data_ptrs=torch.tensor( + [x.data_ptr() for x in v_buffer], dtype=torch.uint64, device=DEVICE + ), + ) + + +def assert_backup_matches_device(host, device_pool, host_indices_host, device_indices): + for layer_id in range(NUM_LAYERS): + torch.testing.assert_close( + host.k_buffer[host_indices_host, layer_id], + device_pool.k_buffer[layer_id][device_indices].cpu(), + ) + torch.testing.assert_close( + host.v_buffer[host_indices_host, layer_id], + device_pool.v_buffer[layer_id][device_indices].cpu(), + ) + + +def assert_load_matches_host(host, device_pool, host_indices_host, load_indices): + for layer_id in range(NUM_LAYERS): + torch.testing.assert_close( + device_pool.k_buffer[layer_id][load_indices], + host.k_buffer[host_indices_host, layer_id].to(DEVICE), + ) + torch.testing.assert_close( + device_pool.v_buffer[layer_id][load_indices], + host.v_buffer[host_indices_host, layer_id].to(DEVICE), + ) + + +@pytest.mark.parametrize("dtype", DTYPES) +def test_asymmetric_mha_kernel_page_first_roundtrip(dtype): + # Covers D2H backup + H2D load through AsymmetricMHATokenToKVPoolHost using + # MiMoV2's real K/V head dims and the real MLA single-buffer kernels. + host = make_host_pool(dtype) + device_pool = make_device_pool(dtype) + + device_pages = torch.tensor([1, 2, 3], dtype=torch.int64) + host_pages = torch.tensor([0, 1, 2], dtype=torch.int64) + load_pages = torch.tensor([4, 5, 6], dtype=torch.int64) + device_indices_host = token_indices_for_pages(device_pages) + host_indices_host = token_indices_for_pages(host_pages) + load_indices_host = token_indices_for_pages(load_pages) + device_indices = device_indices_host.to(DEVICE) + host_indices = host_indices_host.to(DEVICE) + load_indices = load_indices_host.to(DEVICE) + + host.backup_from_device_all_layer( + device_pool, host_indices, device_indices, io_backend="kernel" + ) + torch.cuda.synchronize() + assert_backup_matches_device( + host, device_pool, host_indices_host, device_indices_host + ) + + for layer_id in range(NUM_LAYERS): + device_pool.k_buffer[layer_id].zero_() + device_pool.v_buffer[layer_id].zero_() + host.load_to_device_per_layer( + device_pool, host_indices, load_indices, layer_id, io_backend="kernel" + ) + torch.cuda.synchronize() + assert_load_matches_host(host, device_pool, host_indices_host, load_indices_host) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/models_e2e/test_mimo_v2.py b/test/registered/models_e2e/test_mimo_v2.py index 8f1898a8b..d22019284 100644 --- a/test/registered/models_e2e/test_mimo_v2.py +++ b/test/registered/models_e2e/test_mimo_v2.py @@ -1,5 +1,6 @@ import unittest +from sglang.srt.environ import envs from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.kits.eval_accuracy_kit import GSM8KMixin from sglang.test.server_fixtures.mmmu_fixture import MMMUServerBase @@ -20,6 +21,13 @@ MIMO_V2_OTHER_ARGS = [ "fa3", "--reasoning-parser", "mimo", + "--enable-hierarchical-cache", + "--hicache-ratio", + "1.5", + "--hicache-mem-layout", + "page_first", + "--hicache-io-backend", + "kernel", ] MIMO_V2_MTP_OTHER_ARGS = MIMO_V2_OTHER_ARGS + [ "--speculative-algorithm", @@ -42,6 +50,11 @@ class TestMiMoV2(GSM8KMixin, MMMUServerBase): server_api_key = None other_args = MIMO_V2_MTP_OTHER_ARGS + @classmethod + def setUpClass(cls): + with envs.SGLANG_ENABLE_UNIFIED_RADIX_TREE.override(True): + super().setUpClass() + if __name__ == "__main__": unittest.main() diff --git a/test/registered/models_e2e/test_mimo_v2_flash.py b/test/registered/models_e2e/test_mimo_v2_flash.py index 6615b204e..f2d9a0c38 100644 --- a/test/registered/models_e2e/test_mimo_v2_flash.py +++ b/test/registered/models_e2e/test_mimo_v2_flash.py @@ -1,5 +1,6 @@ import unittest +from sglang.srt.environ import envs from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.kits.eval_accuracy_kit import GSM8KMixin from sglang.test.kits.spec_decoding_kit import SpecDecodingMixin @@ -42,11 +43,23 @@ class TestMiMoV2Flash(GSM8KMixin, SpecDecodingMixin, DefaultServerBase): "--enable-multi-layer-eagle", "--model-loader-extra-config", '{"enable_multithread_load": true,"num_threads": 64}', + "--enable-hierarchical-cache", + "--hicache-ratio", + "1.5", + "--hicache-mem-layout", + "page_first", + "--hicache-io-backend", + "kernel", ] bs_1_speed_thres = 170 accept_length_thres = 3.2 + @classmethod + def setUpClass(cls): + with envs.SGLANG_ENABLE_UNIFIED_RADIX_TREE.override(True): + super().setUpClass() + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/mem_cache/test_asymmetric_mha_pool_host_unit.py b/test/registered/unit/mem_cache/test_asymmetric_mha_pool_host_unit.py new file mode 100644 index 000000000..6c3f0e7ac --- /dev/null +++ b/test/registered/unit/mem_cache/test_asymmetric_mha_pool_host_unit.py @@ -0,0 +1,156 @@ +"""Unit tests for asymmetric MHA host KV pool transfer dispatch.""" + +import unittest +from types import SimpleNamespace +from unittest import mock + +import torch + +from sglang.srt.mem_cache.memory_pool_host import ( + AsymmetricMHATokenToKVPoolHost, + MHATokenToKVPoolHost, + get_mha_host_pool_cls, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +def _make_host(layout: str) -> AsymmetricMHATokenToKVPoolHost: + host = AsymmetricMHATokenToKVPoolHost.__new__(AsymmetricMHATokenToKVPoolHost) + host.layout = layout + host.page_size = 2 + host.layer_num = 3 + host.head_num = 2 + host.head_dim = 4 + host.v_head_dim = 6 + host.dtype = torch.float16 + + if layout == "page_first": + k_dims = (8, host.layer_num, host.head_num, host.head_dim) + v_dims = (8, host.layer_num, host.head_num, host.v_head_dim) + else: + raise ValueError(f"Unsupported test layout: {layout}") + + host.kv_buffer = (torch.empty(k_dims), torch.empty(v_dims)) + return host + + +def _make_device_pool(host: AsymmetricMHATokenToKVPoolHost) -> SimpleNamespace: + size = 8 + k_buffer = [ + torch.empty(size, host.head_num, host.head_dim) for _ in range(host.layer_num) + ] + v_buffer = [ + torch.empty(size, host.head_num, host.v_head_dim) for _ in range(host.layer_num) + ] + return SimpleNamespace( + k_buffer=k_buffer, + v_buffer=v_buffer, + k_data_ptrs=torch.tensor([x.data_ptr() for x in k_buffer], dtype=torch.uint64), + v_data_ptrs=torch.tensor([x.data_ptr() for x in v_buffer], dtype=torch.uint64), + ) + + +class TestAsymmetricMHATokenToKVPoolHost(CustomTestCase): + def test_factory_selects_asymmetric_pool_for_mismatched_kv_dims(self): + symmetric_pool = SimpleNamespace(head_dim=4, v_head_dim=4) + asymmetric_pool = SimpleNamespace(head_dim=4, v_head_dim=6) + + self.assertIs(get_mha_host_pool_cls(symmetric_pool), MHATokenToKVPoolHost) + self.assertIs( + get_mha_host_pool_cls(asymmetric_pool), AsymmetricMHATokenToKVPoolHost + ) + + def test_kernel_load_splits_k_and_v_with_separate_strides(self): + # Dispatch-only test: the CUDA kernel is mocked; this verifies that K and + # V are sent as separate single-buffer calls with their own byte strides. + host = _make_host("page_first") + device_pool = _make_device_pool(host) + host_indices = torch.tensor([0, 1, 2, 3], dtype=torch.int64) + device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64) + + with mock.patch( + "sglang.srt.mem_cache.memory_pool_host.transfer_kv_per_layer_mla_pf_lf", + create=True, + ) as transfer: + host.load_to_device_per_layer( + device_pool, + host_indices, + device_indices, + layer_id=1, + io_backend="kernel", + ) + + self.assertEqual(transfer.call_count, 2) + k_call, v_call = transfer.call_args_list + self.assertIs(k_call.kwargs["src"], host.k_buffer) + self.assertIs(k_call.kwargs["dst"], device_pool.k_buffer[1]) + self.assertEqual(k_call.kwargs["item_size"], 16) + self.assertEqual(k_call.kwargs["src_layout_dim"], 48) + self.assertIs(v_call.kwargs["src"], host.v_buffer) + self.assertIs(v_call.kwargs["dst"], device_pool.v_buffer[1]) + self.assertEqual(v_call.kwargs["item_size"], 24) + self.assertEqual(v_call.kwargs["src_layout_dim"], 72) + + def test_kernel_backup_splits_k_and_v_with_separate_strides(self): + # Dispatch-only test: D2H backup must pass separate K/V layer pointer + # tables so the single-buffer MLA kernel gets the correct stride per side. + host = _make_host("page_first") + device_pool = _make_device_pool(host) + host_indices = torch.tensor([0, 1, 2, 3], dtype=torch.int64) + device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64) + + with mock.patch( + "sglang.srt.mem_cache.memory_pool_host.transfer_kv_all_layer_mla_lf_pf", + create=True, + ) as transfer: + host.backup_from_device_all_layer( + device_pool, host_indices, device_indices, io_backend="kernel" + ) + + self.assertEqual(transfer.call_count, 2) + k_call, v_call = transfer.call_args_list + self.assertIs(k_call.kwargs["src_layers"], device_pool.k_data_ptrs) + self.assertIs(k_call.kwargs["dst"], host.k_buffer) + self.assertEqual(k_call.kwargs["item_size"], 16) + self.assertEqual(k_call.kwargs["dst_layout_dim"], 48) + self.assertIs(v_call.kwargs["src_layers"], device_pool.v_data_ptrs) + self.assertIs(v_call.kwargs["dst"], host.v_buffer) + self.assertEqual(v_call.kwargs["item_size"], 24) + self.assertEqual(v_call.kwargs["dst_layout_dim"], 72) + + def test_direct_load_is_rejected(self): + # Direct single-buffer D2H is not reliable for asymmetric K/V in the + # current sgl-kernel fast path, so the asymmetric host pool is kernel-only. + host = _make_host("page_first") + device_pool = _make_device_pool(host) + host_indices = torch.tensor([0, 1, 2, 3], dtype=torch.int64) + device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64) + + with self.assertRaisesRegex(ValueError, "expected 'kernel'"): + host.load_to_device_per_layer( + device_pool, + host_indices, + device_indices, + layer_id=2, + io_backend="direct", + ) + + def test_direct_backup_is_rejected(self): + # Same restriction for D2H backup: asymmetric MHA uses the kernel path + # until the direct kernel has an explicit safe asymmetric mode. + host = _make_host("page_first") + device_pool = _make_device_pool(host) + host_indices = torch.tensor([0, 1, 2, 3], dtype=torch.int64) + device_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64) + + with self.assertRaisesRegex(ValueError, "expected 'kernel'"): + host.backup_from_device_all_layer( + device_pool, host_indices, device_indices, io_backend="direct" + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py index 04c7cb886..2ec8a57b0 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py @@ -450,16 +450,24 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase): def _init_hicache(self, tree, *, write_policy: str = "write_through"): import sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler as assembler - orig_kv_host_pool = assembler.MHATokenToKVPoolHost + # Wrap the host-pool factory (not MHATokenToKVPoolHost directly) + # because the assembler picks between MHATokenToKVPoolHost and + # AsymmetricMHATokenToKVPoolHost via get_mha_host_pool_cls(device_pool). + orig_get_mha_host_pool_cls = assembler.get_mha_host_pool_cls - def kv_host_pool_wrapper(*args, **kwargs): - kwargs["pin_memory"] = False - return orig_kv_host_pool(*args, **kwargs) + def get_mha_host_pool_cls_wrapper(device_pool): + host_pool_cls = orig_get_mha_host_pool_cls(device_pool) + + def kv_host_pool_wrapper(*args, **kwargs): + kwargs["pin_memory"] = False + return host_pool_cls(*args, **kwargs) + + return kv_host_pool_wrapper patcher = mock.patch.object( assembler, - "MHATokenToKVPoolHost", - side_effect=kv_host_pool_wrapper, + "get_mha_host_pool_cls", + side_effect=get_mha_host_pool_cls_wrapper, ) patcher.start() self.addCleanup(patcher.stop) @@ -2370,12 +2378,20 @@ class UnifiedRadixCacheSuite: ): import sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler as assembler - orig_kv_host_pool = assembler.MHATokenToKVPoolHost + # See _init_hicache: wrap the factory rather than MHATokenToKVPoolHost + # directly so the pin_memory=False override applies to both + # MHATokenToKVPoolHost and AsymmetricMHATokenToKVPoolHost. + orig_get_mha_host_pool_cls = assembler.get_mha_host_pool_cls orig_mamba_host_pool = assembler.MambaPoolHost - def kv_host_pool_wrapper(*args, **kwargs): - kwargs["pin_memory"] = False - return orig_kv_host_pool(*args, **kwargs) + def get_mha_host_pool_cls_wrapper(device_pool): + host_pool_cls = orig_get_mha_host_pool_cls(device_pool) + + def kv_host_pool_wrapper(*args, **kwargs): + kwargs["pin_memory"] = False + return host_pool_cls(*args, **kwargs) + + return kv_host_pool_wrapper def mamba_host_pool_wrapper(*args, **kwargs): kwargs["pin_memory"] = False @@ -2384,8 +2400,8 @@ class UnifiedRadixCacheSuite: patchers = [ mock.patch.object( assembler, - "MHATokenToKVPoolHost", - side_effect=kv_host_pool_wrapper, + "get_mha_host_pool_cls", + side_effect=get_mha_host_pool_cls_wrapper, ), mock.patch.object( assembler,