diff --git a/csrc/cuda/mp_mem_kernels.cu b/csrc/cuda/mp_mem_kernels.cu index 5963b8ab16f..5e0dc7121fd 100644 --- a/csrc/cuda/mp_mem_kernels.cu +++ b/csrc/cuda/mp_mem_kernels.cu @@ -447,8 +447,6 @@ void multi_layer_block_kv_transfer( PageBufferShapeDesc shape_desc, int lmcache_chunk_size, EngineKVFormat engine_kv_format, int skip_prefix_n_blocks) { int head_bytes = shape_desc.hs * shape_desc.element_size; - TORCH_CHECK(head_bytes % sizeof(uint16_t) == 0, "head_size * element_size (", - head_bytes, ") must be divisible by 2 for vectorized access"); if (engine_kv_format == EngineKVFormat::NL_X_NB_BSV_BSS) { // Blocked-scale indexer cache: the per-token fp32 scale must be a whole @@ -464,8 +462,14 @@ void multi_layer_block_kv_transfer( LAUNCH_TEMPLATED(uint4); // 16 bytes per copy } else if (head_bytes % sizeof(uint32_t) == 0) { LAUNCH_TEMPLATED(uint32_t); // 4 bytes per copy + } else if (head_bytes % sizeof(uint16_t) == 0) { + LAUNCH_TEMPLATED(uint16_t); // 2 bytes per copy } else { - LAUNCH_TEMPLATED(uint16_t); // 2 bytes per copy (minimum granularity) + // Opaque model-owned page tails can make a logical row byte-odd (for + // example GLM-5.3 C4's 561-byte FP8 row). The page and object are still + // byte-addressable, so retain correctness with a scalar-byte fallback + // instead of rejecting the whole store or truncating the tail. + LAUNCH_TEMPLATED(uint8_t); // 1 byte per copy (correctness fallback) } } diff --git a/lmcache/integration/vllm/kv_cache_group_edits.py b/lmcache/integration/vllm/kv_cache_group_edits.py index 1db1c8f849f..c8e66b00682 100644 --- a/lmcache/integration/vllm/kv_cache_group_edits.py +++ b/lmcache/integration/vllm/kv_cache_group_edits.py @@ -443,9 +443,13 @@ def apply( groups in the shared pool. head_size = ceil(row / block_size), rounded up to the kernels' - vector alignment, and block_size * head_size may exceed the row - by at most this layer's own page padding - (spec.page_size_bytes), never reaching sibling bytes. + vector alignment. If the page has enough padding, block_size * + head_size may exceed the semantic row while remaining inside this + layer's declared page. Some exact-fit pages cannot be split into + vector-aligned token rows (GLM-5.3-Flash DCP1 is 4096 * 265 bytes). + Those pages use one opaque physical slot instead; the adapter carries + block_size independently as tokens_per_block, so logical scheduling + and cache keys remain unchanged. """ assert isinstance(kv_cache, torch.Tensor), ( "single-layer KV cache must be a torch.Tensor" @@ -475,14 +479,36 @@ def apply( if candidate * block_size * elem <= page_bytes: head_size = candidate break - if head_size == 0: + if head_size != 0: + return kv_cache.as_strided( + (num_blocks, block_size, head_size), + (kv_cache.stride(0), head_size, 1), + ) + + # The semantic row is still a valid opaque page even when it cannot + # be factored into vector-aligned per-token rows. Preserve the whole + # declared page as one physical slot. EngineGroupInfo retains the + # logical block_size, and KVLayerGroupsManager therefore maps the one + # slot back to exactly one engine block ID. + if page_bytes % elem: + raise ValueError( + f"declared Mamba page size {page_bytes} bytes is not aligned " + f"to tensor element size {elem}" + ) + page_elems = page_bytes // elem + if page_elems < row: + raise ValueError( + f"declared Mamba page has {page_elems} elements but the state " + f"row requires {row}" + ) + if page_elems > block_step: raise ValueError( - f"cannot tile a {row}-element state row into {block_size} " - f"aligned tokens within the {page_bytes}-byte page" + f"declared Mamba page has {page_elems} elements but the " + f"physical block stride is only {block_step}" ) return kv_cache.as_strided( - (num_blocks, block_size, head_size), - (kv_cache.stride(0), head_size, 1), + (num_blocks, 1, page_elems), + (kv_cache.stride(0), page_elems, 1), ) @@ -490,12 +516,20 @@ class _PaddedAttentionPageViewEdit(KVCacheGroupEdit): """Canonicalize a padded attention layer as an opaque rank-3 page. Current vLLM HMA layouts expose MLA as ``[B, H=1, N, C]`` and a replicated - DFlash page as e.g. ``[B, H=2, N=16, C=256]``. In both cases the inner page - is tightly packed, while sibling pages create a gap between dim-0 rows. - LMCache's opaque ``[B, N, C]`` format supports that authoritative padded - block stride. Re-factoring all inner page elements over the engine's - logical block size is a zero-copy view and preserves every payload byte; + DFlash page as e.g. ``[B, H=2, N=16, C=256]``. The semantic inner page is + tightly packed, while the declared physical page can also append opaque + model-owned state before sibling pages create the remaining gap between + dim-0 rows. LMCache's opaque ``[B, N, C]`` format supports that + authoritative padded block stride. The complete declared page is exposed + as a zero-copy view, preserving both semantic KV and any opaque page tail; the resulting dimensions are addressing metadata, not semantic K/V axes. + + Ordinary pages retain one physical slot per token. Packed pages whose byte + count cannot be factored over the logical block size (for example GLM-5.3 + NVFP4 KV, whose per-page scale/tail records are not token-uniform) use one + physical slot for the whole page. ``EngineGroupInfo.tokens_per_block`` + independently retains the logical token count, so LMCache's compressed + geometry maps one block ID to one complete opaque page. """ name = "padded-attention-page-view" @@ -520,22 +554,42 @@ def apply( kv_cache: RegisteredKVCache, _layout_hints: LayoutHints, ) -> torch.Tensor: - """Return a padded-stride-preserving opaque ``[B, BS, HS]`` view. + """Return a padded-stride-preserving opaque ``[B, slots, width]`` view. Raises: - ValueError: If one physical page cannot be factored evenly over - the engine's logical block size. + ValueError: If the declared page is not element-aligned, is smaller + than the semantic tensor page, or exceeds the physical dim-0 + stride. """ assert isinstance(kv_cache, torch.Tensor) - page_elems = kv_cache.shape[1:].numel() - if page_elems % spec.block_size: + element_size = kv_cache.element_size() + page_bytes = spec.page_size_bytes + if page_bytes % element_size: + raise ValueError( + f"declared attention page size {page_bytes} bytes is not aligned " + f"to tensor element size {element_size}" + ) + page_elems = page_bytes // element_size + semantic_page_elems = kv_cache.shape[1:].numel() + if page_elems < semantic_page_elems: + raise ValueError( + f"declared attention page has {page_elems} elements but the " + f"semantic tensor page requires {semantic_page_elems}" + ) + if page_elems > kv_cache.stride(0): raise ValueError( - f"a {page_elems}-element attention page cannot be factored " - f"over block_size={spec.block_size}" + f"declared attention page has {page_elems} elements but the " + f"physical block stride is only {kv_cache.stride(0)}" ) - hidden_size = page_elems // spec.block_size + # A packed page need not have a uniform byte width per logical token. + # Treat such a page as one opaque physical slot. The vLLM adapter + # carries spec.block_size separately as tokens_per_block, so the group + # manager derives the correct compression ratio and still consumes one + # engine block ID per logical page. + physical_slots = spec.block_size if page_elems % spec.block_size == 0 else 1 + hidden_size = page_elems // physical_slots return kv_cache.as_strided( - (kv_cache.shape[0], spec.block_size, hidden_size), + (kv_cache.shape[0], physical_slots, hidden_size), (kv_cache.stride(0), hidden_size, 1), ) diff --git a/lmcache/integration/vllm/lmcache_mp_metadata.py b/lmcache/integration/vllm/lmcache_mp_metadata.py index d1237d567fc..abdc8b49a80 100644 --- a/lmcache/integration/vllm/lmcache_mp_metadata.py +++ b/lmcache/integration/vllm/lmcache_mp_metadata.py @@ -566,11 +566,11 @@ def aggregate( ) -> "KVConnectorWorkerMetadata": assert isinstance(other, LMCacheMPWorkerMetadata) merged_requests = dict(self.completed_store_requests) - for k, v in other.completed_store_requests.items(): - merged_requests[k] = merged_requests.get(k, 0) + v + for request_id, count in other.completed_store_requests.items(): + merged_requests[request_id] = merged_requests.get(request_id, 0) + count merged_jobs = dict(self.completed_store_jobs) - for k, v in other.completed_store_jobs.items(): - merged_jobs[k] = merged_jobs.get(k, 0) + v + for job_id, count in other.completed_store_jobs.items(): + merged_jobs[job_id] = merged_jobs.get(job_id, 0) + count return LMCacheMPWorkerMetadata( completed_store_requests=merged_requests, completed_store_jobs=merged_jobs, diff --git a/lmcache/v1/multiprocess/transfer_context/worker_transfer.py b/lmcache/v1/multiprocess/transfer_context/worker_transfer.py index 448bf9fc2c9..dc8f652f82c 100644 --- a/lmcache/v1/multiprocess/transfer_context/worker_transfer.py +++ b/lmcache/v1/multiprocess/transfer_context/worker_transfer.py @@ -670,6 +670,11 @@ def submit_store( RequestType.STORE, [key, instance_id, block_ids, event_ipc_handle], ).to_device_future(device=self._device) + # Multiple incremental stores for one request overwrite the adapter's + # request-keyed event slot. Tie every producer event to its own remote + # future so the exported IPC event remains valid until the sidecar has + # finished waiting on it and reading the corresponding GPU pages. + future.retain_reference(event) self._inflight_store_futures.add(future) return future diff --git a/lmcache/v1/platform/cuda/cumem_ipc.py b/lmcache/v1/platform/cuda/cumem_ipc.py index 49df997f003..e7cddcee654 100644 --- a/lmcache/v1/platform/cuda/cumem_ipc.py +++ b/lmcache/v1/platform/cuda/cumem_ipc.py @@ -1,6 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 """Share vLLM CUDA VMM allocations through same-UID POSIX file descriptors.""" +# Future from __future__ import annotations # Standard diff --git a/tests/v1/multiprocess/test_engine_driven_transfer.py b/tests/v1/multiprocess/test_engine_driven_transfer.py index f7f84dc34b4..aeb1e6d23e0 100644 --- a/tests/v1/multiprocess/test_engine_driven_transfer.py +++ b/tests/v1/multiprocess/test_engine_driven_transfer.py @@ -572,8 +572,8 @@ def test_create_transfer_context_force_lmcache_driven_mode() -> None: assert isinstance(context, LMCacheDrivenTransferContext) -def test_lmcache_driven_preemption_waits_for_remote_store_futures() -> None: - """Handle-path flush waits on remote completion, not the worker device.""" +def test_lmcache_driven_preemption_retains_each_store_event_and_waits() -> None: + """Two stores for one request retain both events until remote completion.""" # First Party from lmcache.v1.multiprocess.transfer_context import ( LMCacheDrivenTransferContext, @@ -582,14 +582,24 @@ def test_lmcache_driven_preemption_waits_for_remote_store_futures() -> None: context = LMCacheDrivenTransferContext() registration_future = MagicMock(name="registration_future") - raw_store_future = MagicMock(name="raw_store_future") - pending = MagicMock(name="pending_store_future") - raw_store_future.to_device_future.return_value = pending + raw_store_futures = [ + MagicMock(name="raw_store_future_1"), + MagicMock(name="raw_store_future_2"), + ] + pending = [ + MagicMock(name="pending_store_future_1"), + MagicMock(name="pending_store_future_2"), + ] + for pending_future in pending: + pending_future.query.return_value = False + for raw_future, pending_future in zip(raw_store_futures, pending, strict=True): + raw_future.to_device_future.return_value = pending_future send_request = MagicMock( - name="send_request", side_effect=[registration_future, raw_store_future] + name="send_request", side_effect=[registration_future, *raw_store_futures] ) event_backend = MagicMock(name="event_backend") - event_backend.export_event.return_value = b"event" + event_backend.export_event.side_effect = [b"event-1", b"event-2"] + events = [MagicMock(name="event_1"), MagicMock(name="event_2")] with ( patch.object( @@ -609,21 +619,27 @@ def test_lmcache_driven_preemption_waits_for_remote_store_futures() -> None: mq_timeout=2.5, send_request=send_request, ) - context.submit_store( - _request_id="request", - key="key", - instance_id=1, - kv_caches={}, - block_ids=[[0]], - event=MagicMock(name="event"), - _blocks_in_chunk=1, - ) + for index, event in enumerate(events): + context.submit_store( + _request_id="request", + key=f"key-{index}", + instance_id=1, + kv_caches={}, + block_ids=[[index]], + event=event, + _blocks_in_chunk=1, + ) + + pending[0].retain_reference.assert_called_once_with(events[0]) + pending[1].retain_reference.assert_called_once_with(events[1]) context.flush_inflight_stores() - pending.result.assert_called_once_with(timeout=2.5) + for pending_future in pending: + pending_future.result.assert_called_once_with(timeout=2.5) context.flush_inflight_stores() - pending.result.assert_called_once() + for pending_future in pending: + pending_future.result.assert_called_once() def test_lmcache_driven_preemption_without_stores_is_noop() -> None: diff --git a/tests/v1/multiprocess/test_ipc_memory_reclaim.py b/tests/v1/multiprocess/test_ipc_memory_reclaim.py index 23659308998..59223e6a7bc 100644 --- a/tests/v1/multiprocess/test_ipc_memory_reclaim.py +++ b/tests/v1/multiprocess/test_ipc_memory_reclaim.py @@ -17,6 +17,7 @@ # Standard # Standard Library from types import SimpleNamespace +from typing import cast from unittest.mock import MagicMock import threading import time @@ -153,7 +154,7 @@ def block_resolve(*_args: object, **_kwargs: object) -> list[list[object]]: raise TimeoutError("active STORE was not released") raise RuntimeError("stop after lifetime check") - module.context.resolve_obj_keys.side_effect = block_resolve + cast(MagicMock, module.context.resolve_obj_keys).side_effect = block_resolve def run_store() -> None: try: diff --git a/tests/v1/test_mp_mem_kernels.py b/tests/v1/test_mp_mem_kernels.py index 2c022fed39c..f13f4f1c5d8 100644 --- a/tests/v1/test_mp_mem_kernels.py +++ b/tests/v1/test_mp_mem_kernels.py @@ -251,6 +251,7 @@ def call_block_kernel( is_mla: bool, tokens_per_object: int, skip_prefix_n_blocks: int = 0, + block_stride_elems: int = 0, ) -> None: device = vllm_tensors[0].device @@ -262,6 +263,7 @@ def call_block_kernel( shape_desc.nh = nh shape_desc.hs = hs shape_desc.element_size = vllm_tensors[0].element_size() + shape_desc.block_stride_elems = block_stride_elems ptrs = [t.data_ptr() for t in vllm_tensors] paged_buffer_ptrs_tensor = torch.tensor(ptrs, dtype=torch.int64, device=device) @@ -402,6 +404,188 @@ def test_block_transfer_roundtrip( ) +def test_block_transfer_roundtrip_packed_nvfp4_page(): + """A 512-token packed NVFP4 page round-trips as one opaque slot.""" + device = torch.device(torch_device_type) + nl, nb, bs, nh, hs = 2, 20, 1, 1, 177_408 + block_stride_elems = 200_000 + slots_per_object = 8 # 4096 logical tokens / 512 tokens per page + total_blocks = slots_per_object + + source_pools = [ + (torch.arange(nb * block_stride_elems, device=device) + layer_idx) + .remainder(251) + .to(torch.uint8) + for layer_idx in range(nl) + ] + target_pools = [ + torch.zeros(nb * block_stride_elems, dtype=torch.uint8, device=device) + for _ in range(nl) + ] + source_vllm = [ + pool.as_strided( + (nb, bs, hs), + (block_stride_elems, hs, 1), + ) + for pool in source_pools + ] + target_vllm = [ + pool.as_strided( + (nb, bs, hs), + (block_stride_elems, hs, 1), + ) + for pool in target_pools + ] + mem_objects = create_memory_objects( + 1, + nl, + slots_per_object, + hs, + 1, + torch.uint8, + device, + ) + + block_ids_d2h = list(range(total_blocks)) + block_ids_h2d = list(range(total_blocks, 2 * total_blocks)) + call_block_kernel( + source_vllm, + mem_objects, + block_ids_d2h, + FMT_MLA, + lmcache_native.TransferDirection.D2H, + nl, + nb, + bs, + nh, + hs, + True, + slots_per_object, + block_stride_elems=block_stride_elems, + ) + call_block_kernel( + target_vllm, + mem_objects, + block_ids_h2d, + FMT_MLA, + lmcache_native.TransferDirection.H2D, + nl, + nb, + bs, + nh, + hs, + True, + slots_per_object, + block_stride_elems=block_stride_elems, + ) + torch_dev.synchronize() + + for src_block, dst_block in zip(block_ids_d2h, block_ids_h2d, strict=True): + for source, target in zip(source_vllm, target_vllm, strict=True): + assert torch.equal(target[dst_block], source[src_block]) + + for target_pool in target_pools: + for block_idx in block_ids_h2d: + padding_start = block_idx * block_stride_elems + hs + padding_end = (block_idx + 1) * block_stride_elems + assert torch.count_nonzero(target_pool[padding_start:padding_end]) == 0 + + +def test_block_transfer_roundtrip_mamba_dcp1_opaque_page(): + """The exact GLM-5.3-Flash DCP1 state-page width round-trips whole.""" + device = torch.device(torch_device_type) + nl, nb, bs, nh, hs = 2, 6, 1, 1, 1_085_440 + block_stride_elems = hs + 64 + slots_per_object = 1 + num_objects = 2 + total_blocks = slots_per_object * num_objects + + source_pools = [ + torch.zeros(nb * block_stride_elems, dtype=torch.uint8, device=device) + for _ in range(nl) + ] + for layer_idx, pool in enumerate(source_pools): + for block_idx in range(nb): + page_start = block_idx * block_stride_elems + page_end = page_start + hs + fill_value = (17 * layer_idx + block_idx + 1) % 251 + pool[page_start:page_end].fill_(fill_value) + pool[page_start : page_start + 32] = torch.arange( + 32, dtype=torch.uint8, device=device + ) + (layer_idx * 32) + target_pools = [ + torch.zeros(nb * block_stride_elems, dtype=torch.uint8, device=device) + for _ in range(nl) + ] + source_vllm = [ + pool.as_strided( + (nb, bs, hs), + (block_stride_elems, hs, 1), + ) + for pool in source_pools + ] + target_vllm = [ + pool.as_strided( + (nb, bs, hs), + (block_stride_elems, hs, 1), + ) + for pool in target_pools + ] + mem_objects = create_memory_objects( + 1, + nl, + slots_per_object, + hs, + num_objects, + torch.uint8, + device, + ) + + block_ids_d2h = list(range(total_blocks)) + block_ids_h2d = list(range(total_blocks, 2 * total_blocks)) + call_block_kernel( + source_vllm, + mem_objects, + block_ids_d2h, + FMT_MLA, + lmcache_native.TransferDirection.D2H, + nl, + nb, + bs, + nh, + hs, + True, + slots_per_object, + block_stride_elems=block_stride_elems, + ) + call_block_kernel( + target_vllm, + mem_objects, + block_ids_h2d, + FMT_MLA, + lmcache_native.TransferDirection.H2D, + nl, + nb, + bs, + nh, + hs, + True, + slots_per_object, + block_stride_elems=block_stride_elems, + ) + torch_dev.synchronize() + + for src_block, dst_block in zip(block_ids_d2h, block_ids_h2d, strict=True): + for source, target in zip(source_vllm, target_vllm, strict=True): + assert torch.equal(target[dst_block], source[src_block]) + + for target_pool in target_pools: + for block_idx in block_ids_h2d: + padding_start = block_idx * block_stride_elems + hs + padding_end = (block_idx + 1) * block_stride_elems + assert torch.count_nonzero(target_pool[padding_start:padding_end]) == 0 + + @pytest.mark.parametrize( "engine_kv_format,nl,nh,hs,is_mla", FORMAT_PARAMS, @@ -649,3 +833,92 @@ def test_block_transfer_roundtrip_large_block(dtype, mem_device): assert torch.equal(src_data[layer_idx], tgt_data[layer_idx]), ( f"Mismatch at block index {i}, layer {layer_idx}" ) + + +def test_block_transfer_roundtrip_byte_odd_padded_mla_page(): + """Opaque byte-odd rows round-trip without touching dim-0 padding.""" + device = torch.device(torch_device_type) + nl, nb, bs, nh, hs = 2, 40, 16, 1, 561 + block_stride_elems = bs * hs + 37 + tokens_per_object = 256 + blocks_per_object = tokens_per_object // bs + num_objects = 1 + total_blocks = blocks_per_object * num_objects + + source_pools = [ + (torch.arange(nb * block_stride_elems, device=device) + layer_idx) + .remainder(251) + .to(torch.uint8) + for layer_idx in range(nl) + ] + target_pools = [ + torch.zeros(nb * block_stride_elems, dtype=torch.uint8, device=device) + for _ in range(nl) + ] + source_vllm = [ + pool.as_strided( + (nb, bs, hs), + (block_stride_elems, hs, 1), + ) + for pool in source_pools + ] + target_vllm = [ + pool.as_strided( + (nb, bs, hs), + (block_stride_elems, hs, 1), + ) + for pool in target_pools + ] + mem_objects = create_memory_objects( + 1, + nl, + tokens_per_object, + hs, + num_objects, + torch.uint8, + device, + ) + + block_ids_d2h = list(range(total_blocks)) + block_ids_h2d = list(range(total_blocks, 2 * total_blocks)) + call_block_kernel( + source_vllm, + mem_objects, + block_ids_d2h, + FMT_MLA, + lmcache_native.TransferDirection.D2H, + nl, + nb, + bs, + nh, + hs, + True, + tokens_per_object, + block_stride_elems=block_stride_elems, + ) + call_block_kernel( + target_vllm, + mem_objects, + block_ids_h2d, + FMT_MLA, + lmcache_native.TransferDirection.H2D, + nl, + nb, + bs, + nh, + hs, + True, + tokens_per_object, + block_stride_elems=block_stride_elems, + ) + torch_dev.synchronize() + + for src_block, dst_block in zip(block_ids_d2h, block_ids_h2d, strict=True): + for source, target in zip(source_vllm, target_vllm, strict=True): + assert torch.equal(target[dst_block], source[src_block]) + + for target_pool in target_pools: + for block_idx in block_ids_h2d: + padding_start = block_idx * block_stride_elems + bs * hs + padding_end = (block_idx + 1) * block_stride_elems + assert torch.count_nonzero(target_pool[padding_start:padding_end]) == 0 diff --git a/tests/v1/test_vllm_dcp_support.py b/tests/v1/test_vllm_dcp_support.py index c5f3dd9a513..f87c64672f6 100644 --- a/tests/v1/test_vllm_dcp_support.py +++ b/tests/v1/test_vllm_dcp_support.py @@ -468,6 +468,86 @@ def test_mamba_unified_view_preserves_blocks_first_pool_stride( ) +@requires_vllm +def test_mamba_dcp1_exact_fit_page_uses_one_opaque_slot(): + """An odd token-row width must not reject an otherwise aligned page. + + GLM-5.3-Flash DCP1 exposes a 1,085,440-byte recurrent page spanning 4096 + logical tokens. A synthetic token row would be 265 bytes, which cannot + satisfy the native transfer vector width without spilling past the page. + The complete page is instead represented by one physical slot while the + group metadata retains its 4096-token logical span. + """ + # Third Party + from vllm.v1.kv_cache_interface import KVCacheConfig, KVCacheGroupSpec + from vllm.v1.kv_cache_interface import MambaSpec as VllmMambaSpec + + # First Party + from lmcache.integration.vllm.kv_cache_group_edits import ( + apply_kv_cache_group_edits, + ) + from lmcache.v1.kv_layer_groups import KVLayerGroupsManager + from lmcache.v1.multiprocess.group_view import EngineGroupInfo + import lmcache.lmcache_native as lmcache_native + + num_blocks = 3 + num_layers = 2 + logical_block_size = 4096 + page_elems = 1_085_440 + block_stride = num_layers * page_elems + pool = torch.zeros(num_blocks * block_stride, dtype=torch.uint8) + layer = pool.as_strided( + (num_blocks, 1, 1, page_elems), + (block_stride, page_elems, page_elems, 1), + storage_offset=page_elems, + ) + layer[1, 0, 0, :32] = torch.arange(32, dtype=torch.uint8) + spec = VllmMambaSpec( + block_size=logical_block_size, + shapes=((page_elems,),), + dtypes=(torch.uint8,), + page_size_padded=page_elems, + mamba_cache_mode="align", + ) + kv_config = KVCacheConfig( + num_blocks=num_blocks, + kv_cache_tensors=[], + kv_cache_groups=[KVCacheGroupSpec(["mamba"], spec)], + kv_cache_layout="BLHNC", + ) + + edited = apply_kv_cache_group_edits( + kv_config, + {"mamba": layer}, + {"kv_layout": "BLHNC"}, + )["mamba"] + + assert isinstance(edited, torch.Tensor) + assert edited.shape == (num_blocks, 1, page_elems) + assert edited.stride() == (block_stride, page_elems, 1) + assert edited.data_ptr() == layer.data_ptr() + assert torch.equal(edited[1, 0], layer[1, 0, 0]) + + manager = KVLayerGroupsManager( + [edited], + engine_kv_formats=[lmcache_native.EngineKVFormat.NL_X_NB_BS_HS], + engine_group_infos=[ + EngineGroupInfo( + engine_group_id=0, + layer_indices=(0,), + tokens_per_block=logical_block_size, + ) + ], + lmcache_tokens_per_chunk=logical_block_size, + ) + group = manager.kernel_groups[0] + assert group.slots_per_block == 1 + assert group.tokens_per_block == logical_block_size + assert group.calculate_slots(logical_block_size) == 1 + assert group.shape_desc.hs == page_elems + assert group.shape_desc.block_stride_elems == block_stride + + @requires_vllm @pytest.mark.parametrize( ("kv_layout", "target_shape", "target_strides"), @@ -622,6 +702,181 @@ def test_padded_attention_page_views_preserve_hma_block_padding( ) +@requires_vllm +def test_padded_attention_page_view_includes_declared_opaque_tail(): + """Model-owned bytes appended to an MLA page must round-trip with KV.""" + # Third Party + from vllm.v1.kv_cache_interface import ( + KVCacheConfig, + KVCacheGroupSpec, + ) + from vllm.v1.kv_cache_interface import MambaSpec as VllmMambaSpec + from vllm.v1.kv_cache_interface import ( + MLAAttentionSpec, + ) + + # First Party + from lmcache.integration.vllm.kv_cache_group_edits import ( + apply_kv_cache_group_edits, + ) + + num_blocks = 3 + semantic_page_elems = 32 + declared_page_elems = 40 + block_stride_elems = 80 + pool = torch.arange(num_blocks * block_stride_elems, dtype=torch.uint8) + attention = pool.as_strided( + (num_blocks, 1, 4, 8), + (block_stride_elems, semantic_page_elems, 8, 1), + ) + mla_spec = MLAAttentionSpec( + block_size=4, + num_kv_heads=1, + head_size=8, + dtype=torch.uint8, + page_size_padded=declared_page_elems, + ) + mamba_spec = VllmMambaSpec( + block_size=4, + shapes=((13,),), + dtypes=(torch.float32,), + page_size_padded=64, + mamba_cache_mode="align", + ) + mamba_pool = torch.zeros(num_blocks * 16, dtype=torch.float32) + mamba = mamba_pool.as_strided( + (num_blocks, 1, 1, 13), + (16, 13, 13, 1), + ) + kv_config = KVCacheConfig( + num_blocks=num_blocks, + kv_cache_tensors=[], + kv_cache_groups=[ + KVCacheGroupSpec(["attention"], mla_spec), + KVCacheGroupSpec(["mamba"], mamba_spec), + ], + kv_cache_layout="NHD", + ) + + edited = apply_kv_cache_group_edits( + kv_config, + {"attention": attention, "mamba": mamba}, + {"kv_layout": "NHD"}, + )["attention"] + + assert isinstance(edited, torch.Tensor) + assert edited.shape == (num_blocks, 4, 10) + assert edited.stride() == (block_stride_elems, 10, 1) + assert edited.data_ptr() == attention.data_ptr() + assert torch.equal( + edited[1].reshape(-1), + pool[block_stride_elems : block_stride_elems + declared_page_elems], + ) + + +@requires_vllm +def test_packed_nvfp4_page_uses_one_opaque_slot_per_logical_block(): + """A non-token-factorable NVFP4 page transfers as one exact block record.""" + # Third Party + from vllm.v1.kv_cache_interface import ( + KVCacheConfig, + KVCacheGroupSpec, + ) + from vllm.v1.kv_cache_interface import MambaSpec as VllmMambaSpec + from vllm.v1.kv_cache_interface import ( + MLAAttentionSpec, + ) + + # First Party + from lmcache.integration.vllm.kv_cache_group_edits import ( + apply_kv_cache_group_edits, + ) + from lmcache.v1.kv_layer_groups import KVLayerGroupsManager + from lmcache.v1.multiprocess.group_view import EngineGroupInfo + import lmcache.lmcache_native as lmcache_native + + num_blocks = 3 + logical_block_size = 512 + semantic_record_bytes = 304 + semantic_page_bytes = logical_block_size * semantic_record_bytes + declared_page_bytes = 177_408 + block_stride_bytes = 200_000 + pool = torch.arange(num_blocks * block_stride_bytes, dtype=torch.int64).to( + torch.uint8 + ) + attention = pool.as_strided( + (num_blocks, 1, logical_block_size, semantic_record_bytes), + ( + block_stride_bytes, + semantic_page_bytes, + semantic_record_bytes, + 1, + ), + ) + mla_spec = MLAAttentionSpec( + block_size=logical_block_size, + num_kv_heads=1, + head_size=semantic_record_bytes, + dtype=torch.uint8, + page_size_padded=declared_page_bytes, + ) + mamba_spec = VllmMambaSpec( + block_size=logical_block_size, + shapes=((13,),), + dtypes=(torch.float32,), + page_size_padded=logical_block_size * torch.float32.itemsize, + mamba_cache_mode="align", + ) + mamba_pool = torch.zeros(num_blocks * logical_block_size, dtype=torch.float32) + mamba = mamba_pool.as_strided( + (num_blocks, 1, 1, 13), + (logical_block_size, 13, 13, 1), + ) + kv_config = KVCacheConfig( + num_blocks=num_blocks, + kv_cache_tensors=[], + kv_cache_groups=[ + KVCacheGroupSpec(["attention"], mla_spec), + KVCacheGroupSpec(["mamba"], mamba_spec), + ], + kv_cache_layout="NHD", + ) + + edited = apply_kv_cache_group_edits( + kv_config, + {"attention": attention, "mamba": mamba}, + {"kv_layout": "NHD"}, + )["attention"] + + assert isinstance(edited, torch.Tensor) + assert edited.shape == (num_blocks, 1, declared_page_bytes) + assert edited.stride() == (block_stride_bytes, declared_page_bytes, 1) + assert edited.data_ptr() == attention.data_ptr() + assert torch.equal( + edited[1].reshape(-1), + pool[block_stride_bytes : block_stride_bytes + declared_page_bytes], + ) + + manager = KVLayerGroupsManager( + [edited], + engine_kv_formats=[lmcache_native.EngineKVFormat.NL_X_NB_BS_HS], + engine_group_infos=[ + EngineGroupInfo( + engine_group_id=0, + layer_indices=(0,), + tokens_per_block=logical_block_size, + ) + ], + lmcache_tokens_per_chunk=4096, + ) + group = manager.kernel_groups[0] + assert group.slots_per_block == 1 + assert group.tokens_per_block == logical_block_size + assert group.calculate_slots(4096) == 8 + assert group.shape_desc.hs == declared_page_bytes + assert group.shape_desc.block_stride_elems == block_stride_bytes + + @requires_vllm @pytest.mark.parametrize("attention_kind", ["mla", "standard"]) def test_subpaged_attention_excludes_incomplete_logical_page_tail( diff --git a/tests/v1/test_vllm_mp_adapter.py b/tests/v1/test_vllm_mp_adapter.py index 92d2184cf0d..14306ebbfec 100644 --- a/tests/v1/test_vllm_mp_adapter.py +++ b/tests/v1/test_vllm_mp_adapter.py @@ -4,7 +4,7 @@ recovery: ``.buildkite/k3_tests/multiprocess/scripts/run-restart-recovery.sh``.""" # Standard -from typing import Callable, ClassVar +from typing import Callable, ClassVar, cast from unittest.mock import MagicMock import gc import os @@ -341,7 +341,7 @@ def failing_factory( kv_caches: dict[str, torch.Tensor], mode: str | MPTransferMode | None, ) -> MagicMock: - context = original_factory(kv_caches, mode) + context = cast(MagicMock, original_factory(kv_caches, mode)) context.register.side_effect = CuMemIPCUnsupportedError("not exportable") return context @@ -385,7 +385,7 @@ def failing_factory( kv_caches: dict[str, torch.Tensor], mode: str | MPTransferMode | None, ) -> MagicMock: - context = original_factory(kv_caches, mode) + context = cast(MagicMock, original_factory(kv_caches, mode)) context.register.side_effect = fail_registration return context @@ -415,7 +415,7 @@ def timeout_factory( kv_caches: dict[str, torch.Tensor], mode: str | MPTransferMode | None, ) -> MagicMock: - context = contexts_factory(kv_caches, mode) + context = cast(MagicMock, contexts_factory(kv_caches, mode)) context.register.side_effect = TimeoutError("server down") return context @@ -449,7 +449,7 @@ def observing_factory( kv_caches: dict[str, torch.Tensor], mode: str | MPTransferMode | None, ) -> MagicMock: - context = contexts_factory(kv_caches, mode) + context = cast(MagicMock, contexts_factory(kv_caches, mode)) context.register.side_effect = observe_registration return context @@ -1443,7 +1443,7 @@ def test_recover_callback_closes_superseded_transfer_ctx( original_factory = adapter_mod.create_transfer_context def failing_register(kv_caches: dict[str, torch.Tensor], mode: str) -> MagicMock: - ctx = original_factory(kv_caches, mode) + ctx = cast(MagicMock, original_factory(kv_caches, mode)) ctx.register.side_effect = TimeoutError("server down") return ctx @@ -1486,7 +1486,7 @@ def delayed_factory( kv_caches: dict[str, torch.Tensor], mode: str | MPTransferMode | None, ) -> MagicMock: - context = original_factory(kv_caches, mode) + context = cast(MagicMock, original_factory(kv_caches, mode)) context.register.side_effect = delayed_register return context @@ -1499,9 +1499,12 @@ def delayed_factory( assert recovery_entered.wait(timeout=5.0) shutdown_done = threading.Event() - shutdown_thread = threading.Thread( - target=lambda: (adapter.shutdown(), shutdown_done.set()) - ) + + def shutdown_adapter() -> None: + adapter.shutdown() + shutdown_done.set() + + shutdown_thread = threading.Thread(target=shutdown_adapter) shutdown_thread.start() assert not shutdown_done.wait(timeout=0.05) assert not any( diff --git a/tests/v1/test_vllm_mp_exact_recurrent_boundaries.py b/tests/v1/test_vllm_mp_exact_recurrent_boundaries.py index ae078fa1ea4..9a76b809cfe 100644 --- a/tests/v1/test_vllm_mp_exact_recurrent_boundaries.py +++ b/tests/v1/test_vllm_mp_exact_recurrent_boundaries.py @@ -14,8 +14,6 @@ from vllm.v1.core.sched.output import SchedulerOutput # noqa: E402 # First Party -import lmcache.integration.vllm.lmcache_mp_connector as connector_mod # noqa: E402 -import lmcache.integration.vllm.lmcache_mp_metadata as metadata_mod # noqa: E402 from lmcache.integration.vllm.lmcache_mp_connector import ( # noqa: E402 LMCacheMPConnector, _is_mamba_group_spec, @@ -26,6 +24,8 @@ LMCacheMPRequestState, LMCacheMPRequestTracker, ) +import lmcache.integration.vllm.lmcache_mp_connector as connector_mod # noqa: E402 +import lmcache.integration.vllm.lmcache_mp_metadata as metadata_mod # noqa: E402 GROUP_TOKENS_PER_BLOCK = [512, 512, 512, 512, 2048, 512] CHUNK_TOKENS = 4096 @@ -64,7 +64,7 @@ def __init__(self) -> None: def touch(self, blocks: list[SimpleNamespace]) -> None: self.touched.append([block.block_id for block in blocks]) - def free_blocks(self, blocks: object) -> None: + def free_blocks(self, blocks: list[SimpleNamespace]) -> None: self.freed.append([block.block_id for block in blocks])