From fc896911b4113d7c084052cb9fd969c0162f3d75 Mon Sep 17 00:00:00 2001 From: xiaozeyu Date: Mon, 27 Jul 2026 21:36:49 +0800 Subject: [PATCH 1/7] [py_connector] rewrite vLLM connector around per-group transfer for hybrid attention Every kv_cache_group (FullAttentionSpec / MambaSpec) becomes an independent transfer unit with its own location spec (tp{rank}_g{group}), block table and data access strategy: token-granular gather/scatter for attention groups, per-block opaque byte copy for mamba/linear state groups. A full-attention model is simply the one-group case. - GroupMeta/TransferGroup abstractions; ReqState tracks per-group block tables - SupportsHMA (request_finished_all_groups) for hybrid memory allocator models - Adapt to vLLM 0.26.0 packed KV layout (num_blocks, heads, block, 2*head_size) - Strided gather/scatter kernel path for padded/strided page layouts - Null-block detection for unmaterialized mamba boundary state --- kv_cache_manager/py_connector/common/types.py | 58 +- .../kernel/batch_gather_scatter_helper.py | 56 +- .../py_connector/vllm/data_transfer.py | 443 +++---- .../py_connector/vllm/metadata.py | 49 +- .../py_connector/vllm/v1_connector.py | 1151 +++++++++-------- 5 files changed, 917 insertions(+), 840 deletions(-) diff --git a/kv_cache_manager/py_connector/common/types.py b/kv_cache_manager/py_connector/common/types.py index bc341e7e7..19a546db2 100644 --- a/kv_cache_manager/py_connector/common/types.py +++ b/kv_cache_manager/py_connector/common/types.py @@ -1,22 +1,54 @@ -from enum import Enum -from typing import Tuple, Dict, Optional, Any +from dataclasses import dataclass, field +from typing import List, Optional -import attrs import torch -@attrs.define(frozen=True) +@dataclass +class TransferGroup: + """One KV cache group = one independent transfer unit. + + A vLLM model exposes ``kv_cache_config.kv_cache_groups``: for pure attention + models there is a single ``FullAttentionSpec`` group; hybrid models expose + several ``MambaSpec`` groups plus one ``FullAttentionSpec`` group. Every + group has its own block table (``block_ids`` is a tuple indexed by group) + in units of its own ``block_size``, and its own storage strategy, so the + connector treats each group as a self-contained transfer unit. + """ + + group_idx: int + # Location spec name registered with the manager, e.g. "tp0_g3". + spec_name: str + # True for attention layers (token-granular strided gather/scatter); + # False for mamba/linear/gdn state layers (per-block opaque byte copy). + is_attention: bool + layer_names: List[str] + # The group's own block table granularity in tokens (spec.block_size). + block_size: int + # Bytes stored per manager block for this whole group (all its layers). + per_block_bytes: int + layer_num: int = 0 + + # --- Attention-only fields (is_attention == True) --- + # int64 tensor of [K0, V0, K1, V1, ...] data ptrs on the compute device. + kvcache_ptr_tensor_gpu: Optional[torch.Tensor] = None + per_token_dim: int = 0 # num_kv_heads * head_size + kernel_block_size: int = 0 # tensor.shape[2] + kv_stride: int = 0 # tensor.stride(0), 0 => contiguous flat layout + block_stride: int = 0 # tensor.stride(1), 0 => contiguous flat layout + + # --- State-only fields (is_attention == False) --- + # Per layer (num_blocks, page_size_bytes) uint8 views into the state storage. + block_view_tensors: List[torch.Tensor] = field(default_factory=list) + page_size_bytes: int = 0 # bytes per block per state layer + + +@dataclass class KVCacheInfo: + """Worker-side registered KV cache description (all groups).""" + tp_rank: int world_size: int - kvcaches: Dict[str, torch.Tensor] - kvcache_ptr_tensor_cpu: torch.Tensor - kvcache_ptr_tensor_gpu: torch.Tensor - all_kvcache_ptr_tensor_gpu: torch.Tensor - layer_num: int - local_token_num: int - per_manager_block_shape: Tuple[int, ...] - per_manager_block_byte_size: int - per_token_per_layer_dim_size: int + groups: List[TransferGroup] device: torch.device dtype: torch.dtype diff --git a/kv_cache_manager/py_connector/kernel/batch_gather_scatter_helper.py b/kv_cache_manager/py_connector/kernel/batch_gather_scatter_helper.py index ebb297505..0f986a633 100644 --- a/kv_cache_manager/py_connector/kernel/batch_gather_scatter_helper.py +++ b/kv_cache_manager/py_connector/kernel/batch_gather_scatter_helper.py @@ -42,8 +42,15 @@ def kv_cache_batch_gather_kernel( NUM_KVCACHE_PTRS: tl.constexpr, # num_layers * kv_count BLOCK_SIZE: tl.constexpr, # 隐藏维度分块大小 DTYPE: tl.constexpr = tl.float16, + kv_stride: tl.constexpr = 0, # stride between K and V (for V pointers) + block_stride: tl.constexpr = 0, # stride between blocks (0 = use flat indexing) + local_block_size: tl.constexpr = 0, # actual block size in tensor (0 = use NUM_TOKENS_PER_BLOCK) ): NUM_DIMS_PER_BLOCK = NUM_TOKENS_PER_BLOCK * NUM_DIMS_PER_TOKEN + + # Determine if using strided layout + USE_STRIDED: tl.constexpr = (block_stride != 0) + EFFECTIVE_LOCAL_BLOCK_SIZE: tl.constexpr = local_block_size if local_block_size > 0 else NUM_TOKENS_PER_BLOCK pid = tl.program_id(0) grid_size = tl.num_programs(0) # 实际grid大小 (如3) @@ -64,6 +71,8 @@ def kv_cache_batch_gather_kernel( # 3. 遍历所有KV缓存指针 (k/v for each layer) for ptr_idx in tl.range(NUM_KVCACHE_PTRS): # 3.1 加载当前层的KV缓存基地址 + # Note: For non-MLA, pointer array is [K0, V0, K1, V1, ...] + # V pointer is already V's base (tensor[1].data_ptr()), no need to add kv_stride kvcache_ptr = tl.load(kv_cache_ptrs_ptr + ptr_idx).to(tl.pointer_type(DTYPE)) # 3.2 计算当前层在dst中的基础偏移 @@ -90,7 +99,16 @@ def kv_cache_batch_gather_kernel( # 从HBM的KV缓存加载数据 # 计算源指针: [BLOCK_SIZE] - src_ptrs = kvcache_ptr + global_token_idx * NUM_DIMS_PER_TOKEN + dim_idx_in_token + if USE_STRIDED: + # Strided layout: convert flat token index to strided offset + # V pointer already includes kv_stride offset, so no need to add it again + kv_block_idx = global_token_idx // EFFECTIVE_LOCAL_BLOCK_SIZE + token_in_kv_block = global_token_idx % EFFECTIVE_LOCAL_BLOCK_SIZE + strided_offset = kv_block_idx * block_stride + token_in_kv_block * NUM_DIMS_PER_TOKEN + src_ptrs = kvcache_ptr + strided_offset + dim_idx_in_token + else: + # Contiguous layout: flat indexing + src_ptrs = kvcache_ptr + global_token_idx * NUM_DIMS_PER_TOKEN + dim_idx_in_token load_mask = mask & token_gather_mask data = tl.load(src_ptrs, mask=load_mask, other=0.0) # 大块连续写入 host memory (PCIe优化) @@ -107,7 +125,10 @@ def batch_gather_kv_caches( dst_block_indices: List[int], # List of dst block indices num_tokens_per_block: int, dim_size_per_token_per_layer: int, - sm_count: int = 3 + sm_count: int = 3, + kv_stride: int = 0, # stride between K and V (for V pointers) + block_stride: int = 0, # stride between blocks (0 = use flat indexing) + local_block_size: int = 0, # actual block size in tensor (0 = use num_tokens_per_block) ): # 配置参数 total_blocks = len(dst_block_indices) @@ -132,6 +153,9 @@ def batch_gather_kv_caches( BLOCK_SIZE=2048, DTYPE=pytorch_dtype_to_triton_dtype(dst_tensor.dtype), num_warps=32, + kv_stride=kv_stride, + block_stride=block_stride, + local_block_size=local_block_size if local_block_size > 0 else num_tokens_per_block, ) # TODO autotune num_warps and BLOCK_SIZE @@ -148,8 +172,15 @@ def kv_cache_batch_scatter_kernel( NUM_KVCACHE_PTRS: tl.constexpr, # num_layers * kv_count BLOCK_SIZE: tl.constexpr, # 隐藏维度分块大小 DTYPE: tl.constexpr = tl.float16, + kv_stride: tl.constexpr = 0, # stride between K and V (for V pointers) + block_stride: tl.constexpr = 0, # stride between blocks (0 = use flat indexing) + local_block_size: tl.constexpr = 0, # actual block size in tensor (0 = use NUM_TOKENS_PER_BLOCK) ): NUM_DIMS_PER_BLOCK = NUM_TOKENS_PER_BLOCK * NUM_DIMS_PER_TOKEN + + # Determine if using strided layout + USE_STRIDED: tl.constexpr = (block_stride != 0) + EFFECTIVE_LOCAL_BLOCK_SIZE: tl.constexpr = local_block_size if local_block_size > 0 else NUM_TOKENS_PER_BLOCK pid = tl.program_id(0) grid_size = tl.num_programs(0) # 实际grid大小 (如3) @@ -170,6 +201,8 @@ def kv_cache_batch_scatter_kernel( # 3. 遍历所有KV缓存指针 (k/v for each layer) for ptr_idx in range(NUM_KVCACHE_PTRS): # 3.1 加载当前层的KV缓存基地址 + # Note: For non-MLA, pointer array is [K0, V0, K1, V1, ...] + # V pointer is already V's base (tensor[1].data_ptr()), no need to add kv_stride kvcache_ptr = tl.load(kv_cache_ptrs_ptr + ptr_idx).to(tl.pointer_type(DTYPE)) # 3.2 计算当前层在src中的基础偏移 @@ -201,7 +234,16 @@ def kv_cache_batch_scatter_kernel( # 向HBM的KV缓存写入数据 # 计算目的指针: [BLOCK_SIZE] - dst_ptrs = kvcache_ptr + global_token_idx * NUM_DIMS_PER_TOKEN + dim_idx_in_token + if USE_STRIDED: + # Strided layout: convert flat token index to strided offset + # V pointer already includes kv_stride offset, so no need to add it again + kv_block_idx = global_token_idx // EFFECTIVE_LOCAL_BLOCK_SIZE + token_in_kv_block = global_token_idx % EFFECTIVE_LOCAL_BLOCK_SIZE + strided_offset = kv_block_idx * block_stride + token_in_kv_block * NUM_DIMS_PER_TOKEN + dst_ptrs = kvcache_ptr + strided_offset + dim_idx_in_token + else: + # Contiguous layout: flat indexing + dst_ptrs = kvcache_ptr + global_token_idx * NUM_DIMS_PER_TOKEN + dim_idx_in_token tl.store(dst_ptrs, data, mask=load_mask) @@ -214,7 +256,10 @@ def batch_scatter_kv_caches( src_block_indices: List[int], # List of src block indices num_tokens_per_block: int, dim_size_per_token_per_layer: int, - sm_count: int = 3 + sm_count: int = 3, + kv_stride: int = 0, # stride between K and V (for V pointers) + block_stride: int = 0, # stride between blocks (0 = use flat indexing) + local_block_size: int = 0, # actual block size in tensor (0 = use num_tokens_per_block) ): # 配置参数 total_blocks = len(src_block_indices) @@ -245,4 +290,7 @@ def batch_scatter_kv_caches( BLOCK_SIZE=2048, DTYPE=pytorch_dtype_to_triton_dtype(src_tensor.dtype), num_warps=32, + kv_stride=kv_stride, + block_stride=block_stride, + local_block_size=local_block_size if local_block_size > 0 else num_tokens_per_block, ) diff --git a/kv_cache_manager/py_connector/vllm/data_transfer.py b/kv_cache_manager/py_connector/vllm/data_transfer.py index 36a16323b..eba0d1e63 100644 --- a/kv_cache_manager/py_connector/vllm/data_transfer.py +++ b/kv_cache_manager/py_connector/vllm/data_transfer.py @@ -1,15 +1,36 @@ -import time -import threading -from concurrent.futures.thread import ThreadPoolExecutor +"""Per-group KV cache transfer between vLLM's paged cache and KVCM storage. + +Each ``TransferGroup`` is an independent transfer unit: + +* Attention groups store token-granular KV; a manager block is gathered/scattered + through the strided Triton kernel (``batch_gather_scatter_helper``) which handles + both the contiguous full-attention layout and the block-strided hybrid layout. +* Mamba/linear/gdn groups store per-block opaque state; a manager block maps to a + single logical block whose raw bytes are copied verbatim. -from typing import Any +The transport itself is layout-agnostic: for every manager block we hand the SDK a +``BlockBuffer`` (a pinned CPU region) and the block's remote URI. Save gathers HBM +-> CPU then ``SaveKvCaches``; load ``LoadKvCaches`` -> CPU then scatters CPU -> HBM. +""" + +import threading +import time +from concurrent.futures import ThreadPoolExecutor import torch from kv_cache_manager.client.pybind import kvcm_py_client +from kv_cache_manager.py_connector.common.tp_coordinator import ( + CoordinateMsgSerializer, TpCoordinatorClient, CoordinateMessage, + SendBlockFinishedEvent, LoadBlockFinishedEvent, +) +from kv_cache_manager.py_connector.common.logger import logger +from kv_cache_manager.py_connector.common.types import KVCacheInfo, TransferGroup +from kv_cache_manager.py_connector.kernel import batch_gather_scatter_helper + def _get_device_module(device=None): - """Return the device module matching the runtime device.""" + """Return the torch device module matching the runtime device.""" if device is not None and hasattr(torch, "get_device_module"): return torch.get_device_module(device) try: @@ -20,25 +41,18 @@ def _get_device_module(device=None): pass return torch.cuda -from kv_cache_manager.py_connector.common.tp_coordinator import CoordinateMsgSerializer, TpCoordinatorClient, \ - CoordinateMessage, SendBlockFinishedEvent, LoadBlockFinishedEvent -from kv_cache_manager.py_connector.common.logger import logger -from kv_cache_manager.py_connector.common.types import KVCacheInfo -from kv_cache_manager.py_connector.kernel import batch_gather_scatter_helper -from kv_cache_manager.py_connector.kernel.gather_scatter_helper import CopyBufferAllocator - class MultiResult: - """多任务结果管理类 - - 用于管理多个异步任务的结果, 当所有任务完成时触发回调 - """ + """Collect the per-block success flags of several async tasks and fire a + callback once every task has reported. Each result is a list[bool] aligned + with the manager blocks the task handled (in submission order).""" + def __init__(self, size: int, callback): - self._size: int = size + self._size = size self._results = [None] * size self._lock = threading.Lock() - self._finished_num: int = 0 - self._finished_callback = callback + self._finished_num = 0 + self._callback = callback def submit_result(self, idx: int, result): with self._lock: @@ -46,247 +60,188 @@ def submit_result(self, idx: int, result): self._results[idx] = result self._finished_num += 1 if self._finished_num == self._size: - self._finished_callback(self._results) + # Flatten in submission order. + flat = [ok for part in self._results for ok in part] + self._callback(flat) class DataTransferManager: - """KVCache数据传输核心类 - - 负责实际的KV缓存保存和加载操作, 包括: - 1. 保存任务(save_task) - 2. 加载任务(load_task) - 3. 回调创建(_create_save_done_callback, _create_load_done_callback) - """ - - def __init__(self, - kvcache_info: KVCacheInfo, - manager_block_size: int, - copy_buffer_allocator: CopyBufferAllocator, - transfer_client: Any, - coordinator_client: TpCoordinatorClient, - extra_config: Any): - """ - 初始化KV数据传输器 - - Args: - kvcache_info: KV缓存信息 - manager_block_size: instance的block_size - copy_buffer_allocator: 复制缓冲区分配器 - transfer_client: 传输客户端 - coordinator_client: 协调器客户端 - extra_config: 额外配置 - """ - self._kvcache_info = kvcache_info + def __init__(self, kvcache_info: KVCacheInfo, manager_block_size: int, + transfer_client, coordinator_client: TpCoordinatorClient, extra_config): + self._info = kvcache_info self._manager_block_size = manager_block_size - self._copy_buffer_allocator = copy_buffer_allocator self._transfer_client = transfer_client self._coordinator_client = coordinator_client self._extra_config = extra_config - self._device_mod = _get_device_module(self._kvcache_info.device) - - # 创建内部线程池执行器 - self._io_executor = self._create_io_executor() - - # 保存和加载流 + self._device = kvcache_info.device + self._device_mod = _get_device_module(self._device) self._save_stream = self._device_mod.Stream() self._load_stream = self._device_mod.Stream() - - def _create_io_executor(self) -> ThreadPoolExecutor: - """创建IO线程池执行器""" - from concurrent.futures import ThreadPoolExecutor - - # 初始化线程池,设置线程名和初始化函数 - def init_worker(): - import torch - self._device_mod.set_device(self._kvcache_info.device) - - return ThreadPoolExecutor( - max_workers=32, - thread_name_prefix="kvcm_io_", - initializer=init_worker - ) - - def submit_task(self, func, *args, **kwargs): - """提交任务到内部线程池 - - Args: - func: 要执行的函数 - *args: 函数参数 - **kwargs: 函数关键字参数 - - Returns: - Future对象 - """ - return self._io_executor.submit(func, *args, **kwargs) - - def load_task(self, multi_result: MultiResult, task_idx, remote_uris, block_token_indices): - """加载任务 - - Args: - multi_result: 多任务结果管理器 - task_idx: 任务索引 - remote_uris: 远程URI列表 - block_token_indices: 块令牌索引列表 - """ - logger.debug("load remote_uris:%s, block_token_indices:%s", remote_uris, block_token_indices) - - copy_buffer_indices = self._copy_buffer_allocator.alloc_buffer_idx_blocking(len(remote_uris)) - copy_buffers = self._copy_buffer_allocator.get_buffer_by_idx(copy_buffer_indices) - - buffers = [] - for copy_buffer in copy_buffers: - buffer = kvcm_py_client.BlockBuffer() - iovs = [] - iov = kvcm_py_client.Iov() - iov.type = kvcm_py_client.MemoryType.CPU - iov.base = copy_buffer.data_ptr() - iov.size = copy_buffer.nbytes - iov.ignore = False - iovs.append(iov) - buffer.iovs = iovs - buffers.append(buffer) - logger.debug("start transfer") - transfer_result = self._transfer_client.LoadKvCaches(remote_uris, buffers) - logger.debug("done transfer,result:%s", transfer_result) - if transfer_result == kvcm_py_client.ClientErrorCode.ER_OK: - with self._device_mod.stream(self._load_stream): - batch_gather_scatter_helper.batch_scatter_kv_caches( - self._kvcache_info.all_kvcache_ptr_tensor_gpu, - self._copy_buffer_allocator._raw_buffer, - block_token_indices, - copy_buffer_indices, - self._manager_block_size, - self._kvcache_info.per_token_per_layer_dim_size, - ) - - copy_done_event = self._device_mod.Event() - copy_done_event.record(self._load_stream) - copy_done_event.synchronize() - - logger.debug("done scatter") - else: - logger.warning("load task failed, remote_uris:%s, block_token_indices:%s, transfer_result:%s", - remote_uris, - block_token_indices, transfer_result) - self._copy_buffer_allocator.free_buffer(copy_buffer_indices) - multi_result.submit_result(task_idx, [transfer_result] * len(remote_uris)) - - def create_load_done_callback(self, req_id, tp_rank, epoch, local_block_ids): - """创建加载完成回调函数 - - Args: - req_id: 请求ID - tp_rank: TP rank - epoch - local_block_ids: 本地块ID列表 - - Returns: - 回调函数 - """ - def generate_message(task_results): - failed_block_idxs = [] - idx = 0 - for task_result in task_results: - for block_result in task_result: - if block_result != kvcm_py_client.ClientErrorCode.ER_OK: - failed_block_idxs.append(local_block_ids[idx]) - idx += 1 - - msg = CoordinateMessage( - time.time(), - LoadBlockFinishedEvent(request_id=req_id, tp_rank=tp_rank, - epoch=epoch, failed_block_idxs=failed_block_idxs) - ) - self._coordinator_client.send(CoordinateMsgSerializer.dumps(msg)) - - return generate_message - def save_task(self, multi_result: MultiResult, task_idx, remote_uris, block_token_indices, - kvcache_ready_event): - """保存任务 - - Args: - multi_result: 多任务结果管理器 - task_idx: 任务索引 - remote_uris: 远程URI列表 - block_token_indices: 块令牌索引列表 - kvcache_ready_event: KV缓存就绪事件 - """ - logger.debug("save remote_uris:%s, block_token_indices:%s", remote_uris, block_token_indices) - - with self._device_mod.stream(self._save_stream): - kvcache_ready_event.wait() - copy_buffer_indices = self._copy_buffer_allocator.alloc_buffer_idx_blocking(len(remote_uris)) - batch_gather_scatter_helper.batch_gather_kv_caches( - self._kvcache_info.all_kvcache_ptr_tensor_gpu, - self._copy_buffer_allocator._raw_buffer, - block_token_indices, - copy_buffer_indices, - self._manager_block_size, - self._kvcache_info.per_token_per_layer_dim_size, - ) - copy_done_event = self._device_mod.Event() - copy_done_event.record(self._save_stream) + def _init_worker(): + self._device_mod.set_device(self._device) - copy_done_event.synchronize() + self._io_executor = ThreadPoolExecutor( + max_workers=32, thread_name_prefix="kvcm_io_", initializer=_init_worker) - logger.debug("done gather") + def submit_task(self, func, *args, **kwargs): + return self._io_executor.submit(func, *args, **kwargs) - copy_buffers = self._copy_buffer_allocator.get_buffer_by_idx(copy_buffer_indices) + # ------------------------------------------------------------------ # + # BlockBuffer helper + # ------------------------------------------------------------------ # + @staticmethod + def _make_block_buffers(base_ptr: int, per_block_bytes: int, count: int): buffers = [] - for copy_buffer in copy_buffers: - buffer = kvcm_py_client.BlockBuffer() - iovs = [] + for i in range(count): + buf = kvcm_py_client.BlockBuffer() iov = kvcm_py_client.Iov() iov.type = kvcm_py_client.MemoryType.CPU - iov.base = copy_buffer.data_ptr() - iov.size = copy_buffer.nbytes + iov.base = base_ptr + i * per_block_bytes + iov.size = per_block_bytes iov.ignore = False - iovs.append(iov) - buffer.iovs = iovs - buffers.append(buffer) - logger.debug("start transfer") - - transfer_result = self._transfer_client.SaveKvCaches(remote_uris, buffers) - logger.debug("done transfer,result:%s", transfer_result) - if transfer_result[0] != kvcm_py_client.ClientErrorCode.ER_OK: - logger.warning("save task failed, remote_uris:%s, block_token_indices:%s, transfer_result:%s", remote_uris, - block_token_indices, transfer_result) - - self._copy_buffer_allocator.free_buffer(copy_buffer_indices) - # TODO: submit uri when enable local alloc - multi_result.submit_result(task_idx, [transfer_result[0]] * len(remote_uris)) - - def create_save_done_callback(self, req_id, tp_rank, write_session_id): - """创建保存完成回调函数 - - Args: - req_id: 请求ID - tp_rank: TP rank - write_session_id: 写入会话ID - - Returns: - 回调函数 + buf.iovs = [iov] + buffers.append(buf) + return buffers + + # ------------------------------------------------------------------ # + # Save + # ------------------------------------------------------------------ # + def save_task(self, multi_result: MultiResult, task_idx, group: TransferGroup, + remote_uris, block_token_indices, block_ids, ready_event): + """Gather one group's manager blocks from HBM and save them. + + block_token_indices: attention -> list[list[int]] flat token slots per block. + block_ids: state -> list[int] block id per manager block; + id 0 is vLLM's null block: the boundary state was + never materialized, so that block cannot be saved. """ - def generate_message(task_results): - is_successes = [] - # TODO: report uri when enable local alloc - # remote_uris = [] - for task_result in task_results: - for block_result in task_result: - if block_result != kvcm_py_client.ClientErrorCode.ER_OK: - is_successes.append(False) - # remote_uris.append(None) - else: - is_successes.append(True) - # remote_uris.extend(future_result[1]) - - msg = CoordinateMessage( - time.time(), - SendBlockFinishedEvent(request_id=req_id, tp_rank=tp_rank, - write_session_id=write_session_id, - is_success_list=is_successes) - ) + n = len(remote_uris) + if group.is_attention: + valid = list(range(n)) + else: + valid = [i for i in range(n) if block_ids[i] != 0] + if len(valid) < n: + logger.warning("save group %s: %d/%d blocks have no materialized " + "state, failing them", group.spec_name, n - len(valid), n) + ok_mask = [False] * n + if valid: + cpu_buffer = torch.empty(len(valid) * group.per_block_bytes, dtype=torch.uint8, + device="cpu", pin_memory=True) + with self._device_mod.stream(self._save_stream): + ready_event.wait() + gpu_buffer = torch.empty(len(valid) * group.per_block_bytes, + dtype=torch.uint8, device=self._device) + if group.is_attention: + view = gpu_buffer.view(self._info.dtype).view( + len(valid), group.layer_num, + self._manager_block_size, group.per_token_dim) + batch_gather_scatter_helper.batch_gather_kv_caches( + group.kvcache_ptr_tensor_gpu, view, block_token_indices, + list(range(len(valid))), self._manager_block_size, + group.per_token_dim, + kv_stride=group.kv_stride, block_stride=group.block_stride, + local_block_size=group.kernel_block_size) + else: + for out_i, i in enumerate(valid): + for layer_idx in range(group.layer_num): + dst = (out_i * group.layer_num + layer_idx) * group.page_size_bytes + gpu_buffer[dst:dst + group.page_size_bytes].copy_( + group.block_view_tensors[layer_idx][block_ids[i]]) + cpu_buffer.copy_(gpu_buffer, non_blocking=True) + done = self._device_mod.Event() + done.record(self._save_stream) + done.synchronize() + + buffers = self._make_block_buffers( + cpu_buffer.data_ptr(), group.per_block_bytes, len(valid)) + uris = [remote_uris[i] for i in valid] + result = self._transfer_client.SaveKvCaches(uris, buffers) + ok = (result[0] == kvcm_py_client.ClientErrorCode.ER_OK) + if not ok: + logger.warning("save task failed group=%s uris=%d result=%s", + group.spec_name, len(uris), result) + for i in valid: + ok_mask[i] = ok + multi_result.submit_result(task_idx, ok_mask) + + def create_save_done_callback(self, req_id, tp_rank, write_session_id, num_blocks): + """block success = AND across all groups. task results are ordered + group0[blocks], group1[blocks], ... so we AND stride-wise.""" + def cb(flat): + is_success = [True] * num_blocks + for i, ok in enumerate(flat): + is_success[i % num_blocks] = is_success[i % num_blocks] and ok + msg = CoordinateMessage(time.time(), SendBlockFinishedEvent( + request_id=req_id, tp_rank=tp_rank, + write_session_id=write_session_id, is_success_list=is_success)) self._coordinator_client.send(CoordinateMsgSerializer.dumps(msg)) - - return generate_message + return cb + + # ------------------------------------------------------------------ # + # Load + # ------------------------------------------------------------------ # + def load_task(self, multi_result: MultiResult, task_idx, group: TransferGroup, + remote_uris, block_token_indices, block_ids): + n = len(remote_uris) + if not group.is_attention and any(b == 0 for b in block_ids): + # Null block: nowhere to scatter the state. Should not happen for + # loads (vLLM allocates real blocks for external tokens). + logger.warning("load group %s: null block in targets, failing task", + group.spec_name) + multi_result.submit_result(task_idx, [False] * n) + return + cpu_buffer = torch.empty(n * group.per_block_bytes, dtype=torch.uint8, + device="cpu", pin_memory=True) + buffers = self._make_block_buffers(cpu_buffer.data_ptr(), group.per_block_bytes, n) + result = self._transfer_client.LoadKvCaches(remote_uris, buffers) + ok = (result == kvcm_py_client.ClientErrorCode.ER_OK) + if ok: + with self._device_mod.stream(self._load_stream): + gpu_buffer = cpu_buffer.to(self._device, non_blocking=True) + if group.is_attention: + view = gpu_buffer.view(self._info.dtype).view( + n, group.layer_num, self._manager_block_size, group.per_token_dim) + batch_gather_scatter_helper.batch_scatter_kv_caches( + group.kvcache_ptr_tensor_gpu, view, block_token_indices, + list(range(n)), self._manager_block_size, group.per_token_dim, + kv_stride=group.kv_stride, block_stride=group.block_stride, + local_block_size=group.kernel_block_size) + else: + for i, block_id in enumerate(block_ids): + for layer_idx in range(group.layer_num): + src = (i * group.layer_num + layer_idx) * group.page_size_bytes + group.block_view_tensors[layer_idx][block_id].copy_( + gpu_buffer[src:src + group.page_size_bytes]) + done = self._device_mod.Event() + done.record(self._load_stream) + done.synchronize() + else: + logger.warning("load task failed group=%s uris=%d result=%s", + group.spec_name, n, result) + multi_result.submit_result(task_idx, [ok] * n) + + def create_load_done_callback(self, req_id, tp_rank, epoch, block_ids, num_blocks, + report_failures=True): + """A manager block is loaded only if every group succeeded for it. + + block_ids is the block table used to report vLLM-visible invalid block + ids. vLLM's invalid-block recovery only understands single-group block + tables, so multi-group (hybrid) connectors pass report_failures=False + and rely on request rescheduling instead.""" + def cb(flat): + merged = [True] * num_blocks + for i, ok in enumerate(flat): + merged[i % num_blocks] = merged[i % num_blocks] and ok + failed = [] + if report_failures: + failed = [block_ids[i] for i in range(min(num_blocks, len(block_ids))) + if not merged[i]] + elif not all(merged): + logger.warning("load failed for %d/%d blocks of req %s (hybrid: " + "not reporting invalid block ids)", + merged.count(False), num_blocks, req_id) + msg = CoordinateMessage(time.time(), LoadBlockFinishedEvent( + request_id=req_id, tp_rank=tp_rank, epoch=epoch, failed_block_idxs=failed)) + self._coordinator_client.send(CoordinateMsgSerializer.dumps(msg)) + return cb diff --git a/kv_cache_manager/py_connector/vllm/metadata.py b/kv_cache_manager/py_connector/vllm/metadata.py index d89d872b8..b63e1bf8c 100644 --- a/kv_cache_manager/py_connector/vllm/metadata.py +++ b/kv_cache_manager/py_connector/vllm/metadata.py @@ -1,56 +1,58 @@ from dataclasses import dataclass, field +from typing import List + from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorMetadata @dataclass class SaveRequest: req_id: str - target_locations: list[dict] - manager_block_idxes: list + # CacheLocation dicts returned by the manager (one per manager block to save), + # each carrying location_specs for every registered spec name. + target_locations: List[dict] + # Manager block indices (into the request's token stream) being saved. + manager_block_idxes: List[int] write_session_id: str -@dataclass() +@dataclass class LoadRequest: req_id: str - manager_block_idxes: list - need_load_locations: list[dict] - local_block_ids: list = field(default_factory=list) + manager_block_idxes: List[int] + need_load_locations: List[dict] + # Per-group block tables: all_block_ids[group_idx] is the list of local block + # ids for that kv_cache_group. Length 1 for pure-attention models. + all_block_ids: List[List[int]] = field(default_factory=list) -@dataclass() +@dataclass class FinishRequest: req_id: str @dataclass class ReqStateToWorker: - """发送给工作节点的请求状态数据结构""" + """Scheduler -> worker per-request state delta.""" req_id: str has_saved_block_num: int new_tokens_ids: list = field(default_factory=list) - new_local_block_ids: list = field(default_factory=list) + # Per-group new local block ids (indexed by kv_cache_group). + new_block_ids_per_group: List[List[int]] = field(default_factory=list) resumed_from_preemption: bool = False is_delta: bool = True + @dataclass class TairKvCacheConnectorMetadata(KVConnectorMetadata): - """TairKvCacheConnector的元数据类,用于在调度器和工作节点之间传递状态""" - requests: list[ReqStateToWorker] + """Scheduler -> worker metadata for one engine step.""" def __init__(self, epoch: int): - """ - 初始化元数据 - - Args: - epoch: 当前epoch编号 - """ self.epoch = epoch - self.requests: list[ReqStateToWorker] = [] - self.to_load_requests: list[LoadRequest] = [] - self.to_save_requests: list[SaveRequest] = [] - self.to_finish_requests: list[FinishRequest] = [] + self.requests: List[ReqStateToWorker] = [] + self.to_load_requests: List[LoadRequest] = [] + self.to_save_requests: List[SaveRequest] = [] + self.to_finish_requests: List[FinishRequest] = [] def add_req_state_to_worker(self, request: ReqStateToWorker): self.requests.append(request) @@ -65,5 +67,6 @@ def add_finish_request(self, finish_request: FinishRequest): self.to_finish_requests.append(finish_request) def __repr__(self): - return f"TairKvCacheConnectorMetadata(requests={self.requests})" - + return (f"TairKvCacheConnectorMetadata(epoch={self.epoch}, " + f"requests={len(self.requests)}, load={len(self.to_load_requests)}, " + f"save={len(self.to_save_requests)}, finish={len(self.to_finish_requests)})") diff --git a/kv_cache_manager/py_connector/vllm/v1_connector.py b/kv_cache_manager/py_connector/vllm/v1_connector.py index 3a0de4d98..d4acd634d 100644 --- a/kv_cache_manager/py_connector/vllm/v1_connector.py +++ b/kv_cache_manager/py_connector/vllm/v1_connector.py @@ -1,13 +1,32 @@ +"""KVCM vLLM connector (v1), built around per-group transfer. + +vLLM models expose one or more ``kv_cache_groups`` (``KVCacheConfig``): + +* Pure-attention models: a single ``FullAttentionSpec`` group. +* Hybrid models (e.g. Qwen3.5): several ``MambaSpec`` groups plus one (or more) + ``FullAttentionSpec`` group. With ``mamba_cache_mode="align"`` every group has + its own block table (``block_ids`` is a tuple indexed by group) but all groups + share the scheduler block size. + +The connector treats every group as an independent transfer unit with its own +KVCM location spec (``tp{rank}_g{group}``), its own block table and its own data +access strategy (token-granular gather/scatter for attention, per-block byte +copy for mamba state). There is no separate "hybrid path": a full-attention +model is simply the one-group case. + +A manager block covers the same token range in every group, so one KVCM cache +key (hashed from token ids) owns the location specs of all groups of all ranks. +""" + import copy import json import math import time import typing -import inspect import threading from dataclasses import dataclass, field -from typing import Any, Optional, List, Dict, Tuple +from typing import Any, List, Optional, Tuple from concurrent.futures import ThreadPoolExecutor from kv_cache_manager.client.pybind import kvcm_py_client @@ -20,6 +39,7 @@ KVConnectorBase_V1, KVConnectorMetadata, KVConnectorRole, + SupportsHMA, ) try: @@ -30,6 +50,7 @@ # vllm <= v0.11.0 from vllm.utils import get_kv_cache_torch_dtype, get_ip +from vllm.v1.kv_cache_interface import FullAttentionSpec, MambaSpec from vllm.v1.core.sched.output import SchedulerOutput from vllm.v1.outputs import KVConnectorOutput @@ -39,8 +60,7 @@ from kv_cache_manager.py_connector.common.logger import logger, configure_log_level from kv_cache_manager.py_connector.common._version_info import FULL_VERSION, GIT_COMMIT, BUILD_TIME -from kv_cache_manager.py_connector.common.types import KVCacheInfo -from kv_cache_manager.py_connector.kernel.gather_scatter_helper import CopyBufferAllocator +from kv_cache_manager.py_connector.common.types import KVCacheInfo, TransferGroup from kv_cache_manager.py_connector.vllm.metadata import SaveRequest, LoadRequest, FinishRequest, ReqStateToWorker, \ TairKvCacheConnectorMetadata from kv_cache_manager.py_connector.vllm.config import TairKvCacheConnectorExtraConfig @@ -52,197 +72,190 @@ from vllm.attention import AttentionMetadata from vllm.v1.request import Request from vllm.v1.core.kv_cache_manager import KVCacheBlocks + from vllm.v1.kv_cache_interface import KVCacheConfig + + +@dataclass +class GroupMeta: + """Static description of one kv_cache_group, derived from KVCacheConfig. + + Available in both scheduler and worker roles (before tensors exist).""" + + group_idx: int + is_attention: bool + layer_names: List[str] + # The group's block table granularity in tokens (spec.block_size). + block_size: int + # Bytes stored per manager block for the whole group. + per_block_bytes: int + # Mamba only: bytes per block per layer (page_size_bytes of the spec). + page_size_bytes: int = 0 @dataclass class ReqState: - """请求状态类,跟踪单个请求的状态信息""" + """Tracks one request. Lives in the scheduler and (mirrored) in workers.""" - # TODO: split this class to ReqStateInScheduler and ReqStateInWorker req_id: str - token_ids: list[int] - local_block_ids: list[int] + token_ids: list + # Per kv_cache_group block table (same length across groups). + block_ids_per_group: List[List[int]] has_saved_block_num: int local_matched_token_num: int remote_matched_token_num: int - # vllm_request only avail in scheduler + # vllm_request only available in scheduler vllm_request: Optional["Request"] - # scheduled_saving_count, sent_saving_count, need_report_after_saving_finished: - # not sync between scheduler and worker and have different meaning - # only available in scheduler and tp0 worker + # Saving progress counters; only meaningful in scheduler and tp0 worker. scheduled_saving_count: int = 0 sent_saving_count: int = 0 need_report_after_saving_finished: bool = False + @property + def num_allocated_blocks(self) -> int: + if not self.block_ids_per_group: + return 0 + return min(len(b) for b in self.block_ids_per_group) + @staticmethod - def create_from_delta(req_state_delta: 'ReqStateToWorker'): - """从ReqStateToWorker创建ReqState实例""" + def create_from_delta(delta: "ReqStateToWorker") -> "ReqState": return ReqState( - req_id=req_state_delta.req_id, - token_ids=req_state_delta.new_tokens_ids, - local_block_ids=req_state_delta.new_local_block_ids, - has_saved_block_num=req_state_delta.has_saved_block_num, + req_id=delta.req_id, + token_ids=list(delta.new_tokens_ids), + block_ids_per_group=[list(b) for b in delta.new_block_ids_per_group], + has_saved_block_num=delta.has_saved_block_num, local_matched_token_num=0, remote_matched_token_num=0, - vllm_request=None + vllm_request=None, ) - def update_from_delta(self, req_state_delta: 'ReqStateToWorker'): - """使用ReqStateToWorker更新当前状态""" - self.token_ids.extend(req_state_delta.new_tokens_ids) - - if req_state_delta.resumed_from_preemption: - self.local_block_ids = req_state_delta.new_local_block_ids + def update_from_delta(self, delta: "ReqStateToWorker"): + self.token_ids.extend(delta.new_tokens_ids) + if not delta.new_block_ids_per_group: + return + if delta.resumed_from_preemption: + self.block_ids_per_group = [list(b) for b in delta.new_block_ids_per_group] else: - self.local_block_ids.extend(req_state_delta.new_local_block_ids) + if not self.block_ids_per_group: + self.block_ids_per_group = [[] for _ in delta.new_block_ids_per_group] + for group_ids, new_ids in zip(self.block_ids_per_group, delta.new_block_ids_per_group): + group_ids.extend(new_ids) -@dataclass -class TransferTaskArgs: - blocks_idx: List[List[int]] = field(default_factory=list) - remote_uris: List[str] = field(default_factory=list) - - -class TairKvCacheConnector(KVConnectorBase_V1): - def _tp_rank_to_spec_name(self, tp_rank: int) -> str: - """Convert TP rank to location spec name.""" - return f"tp{tp_rank}" - - def __init__(self, - vllm_config: "VllmConfig", - role: KVConnectorRole, - kv_cache_config: Optional["KVCacheConfig"] = None, - ): - - init_params = inspect.signature(KVConnectorBase_V1.__init__).parameters - if len(init_params) == 3: - # vllm <= 0.11.0 - super().__init__(vllm_config, role) - else: - # vllm >= 0.11.1 - super().__init__(vllm_config, role, kv_cache_config) +class TairKvCacheConnector(KVConnectorBase_V1, SupportsHMA): - logger.warning("KVCM vllm connector version: %s (commit: %s, build: %s)", FULL_VERSION, GIT_COMMIT, BUILD_TIME) + # ------------------------------------------------------------------ # + # Init / registration + # ------------------------------------------------------------------ # + def __init__(self, vllm_config: "VllmConfig", role: KVConnectorRole, + kv_cache_config: Optional["KVCacheConfig"] = None): + super().__init__(vllm_config, role, kv_cache_config) + assert kv_cache_config is not None, \ + "TairKvCacheConnector requires vLLM to pass kv_cache_config (vllm >= 0.11.1)" - connector_extra_config = vllm_config.kv_transfer_config.kv_connector_extra_config - self._extra_config = TairKvCacheConnectorExtraConfig(connector_extra_config) + logger.warning("KVCM vllm connector version: %s (commit: %s, build: %s)", + FULL_VERSION, GIT_COMMIT, BUILD_TIME) - # Apply log level with priority: env var > startup param > default + self._extra_config = TairKvCacheConnectorExtraConfig( + vllm_config.kv_transfer_config.kv_connector_extra_config) configure_log_level(self._extra_config.log_level) - self._kv_caches: Optional[dict[str, torch.Tensor]] = None - self._local_block_size = vllm_config.cache_config.block_size - model_config = vllm_config.model_config + assert vllm_config.parallel_config.pipeline_parallel_size == 1 + if getattr(model_config, "use_mla", False): + raise NotImplementedError("MLA models are not supported by TairKvCacheConnector") - self._use_mla = (hasattr(model_config, "use_mla") and - isinstance(model_config.use_mla, bool) and - model_config.use_mla) - - manager_block_size = self._local_block_size + self._vllm_block_size = vllm_config.cache_config.block_size + self._tp_size = vllm_config.parallel_config.tensor_parallel_size + self._kv_dtype = get_kv_cache_torch_dtype( + vllm_config.cache_config.cache_dtype, model_config.dtype) + + # Manager block size: attention KV is token-granular and can be re-blocked, + # but mamba state exists once per scheduler block, so hybrid models must + # keep manager block == scheduler block. + manager_block_size = self._vllm_block_size + self._has_state_groups = any( + isinstance(g.kv_cache_spec, MambaSpec) for g in kv_cache_config.kv_cache_groups) if self._extra_config.preferred_block_size != 0: - manager_block_size = self._extra_config.preferred_block_size + if self._has_state_groups: + if self._extra_config.preferred_block_size != self._vllm_block_size: + logger.warning( + "preferred_block_size=%d ignored for hybrid model: mamba state is " + "per scheduler block (%d)", self._extra_config.preferred_block_size, + self._vllm_block_size) + else: + manager_block_size = self._extra_config.preferred_block_size + self._manager_block_size = manager_block_size - self._tp_size = vllm_config.parallel_config.tensor_parallel_size - kv_dtype = get_kv_cache_torch_dtype(vllm_config.cache_config.cache_dtype, model_config.dtype) - num_layer = model_config.get_num_layers(vllm_config.parallel_config) - per_tp_rank_kv_head_num = model_config.get_num_kv_heads(vllm_config.parallel_config) - head_size = model_config.get_head_size() - per_manager_location_spec_shape = [num_layer, 1 if self._use_mla else 2, manager_block_size, - per_tp_rank_kv_head_num, - head_size] + self._group_metas = self._parse_groups(kv_cache_config) + self._num_groups = len(self._group_metas) - assert vllm_config.parallel_config.pipeline_parallel_size == 1 deployment = { "model_name": model_config.served_model_name, - "dtype": str(kv_dtype)[6:], # remove "torch." - "use_mla": self._use_mla, - "tp_size": vllm_config.parallel_config.tensor_parallel_size, + "dtype": str(self._kv_dtype)[6:], # strip "torch." + "use_mla": False, + "tp_size": self._tp_size, "dp_size": vllm_config.parallel_config.data_parallel_size, "pp_size": vllm_config.parallel_config.pipeline_parallel_size, } - logger.info(deployment) + logger.info("deployment: %s, groups: %s", deployment, self._group_metas) - self._manager_client = KvCacheManagerClient.from_connector_config( - vars(self._extra_config) - ) - self._manager_block_size = manager_block_size + self._manager_client = KvCacheManagerClient.from_connector_config(vars(self._extra_config)) self._alive_requests: dict[str, ReqState] = {} self._waiting_to_load_requests: List[LoadRequest] = [] self._waiting_to_save_requests_lock = threading.Lock() self._waiting_to_save_requests: List[SaveRequest] = [] self._waiting_to_finish_requests: List[FinishRequest] = [] - self._canceled_save_request_ids_lock = threading.Lock() self._canceled_save_request_ids: List[str] = [] - # TODO: add coordinator host auto detection, maybe use data parallel host - # TODO: add DP support self._host_ip = get_ip() port = self._extra_config.coordinator_base_port register_response = self._manager_client.register_instance({ - "trace_id": "trace_trace", + "trace_id": "register_%s" % self._extra_config.instance_id, "instance_group": self._extra_config.instance_group, "instance_id": self._extra_config.instance_id, "model_deployment": deployment, "block_size": manager_block_size, - "location_spec_infos": [{ - "name": self._tp_rank_to_spec_name(rank), - "size": math.prod(per_manager_location_spec_shape) * kv_dtype.itemsize - } for rank in range(self._tp_size)], + "location_spec_infos": [ + {"name": self._spec_name(rank, meta.group_idx), "size": meta.per_block_bytes} + for rank in range(self._tp_size) for meta in self._group_metas + ], }) - # TODO: check conflict and update - self._iov_size = math.prod( - per_manager_location_spec_shape) * kv_dtype.itemsize * self._extra_config.hf3fs_concurrent_io_block_count + + max_group_bytes = max(m.per_block_bytes for m in self._group_metas) + self._iov_size = max_group_bytes * self._extra_config.hf3fs_concurrent_io_block_count if role == KVConnectorRole.SCHEDULER: self._epoch = 0 self._coordinator_client = TpCoordinatorClient(self._host_ip, port) self._http_executor = ThreadPoolExecutor(max_workers=4, thread_name_prefix="kvcm_http_") - self._location_query_manager = LocationQueryManager(self._manager_client, self._http_executor, - self._extra_config.instance_id, - self._extra_config.async_get_cache_location) - + self._location_query_manager = LocationQueryManager( + self._manager_client, self._http_executor, self._extra_config.instance_id, + self._extra_config.async_get_cache_location) logger.warning( - "TairKvCacheConnector in scheduler inited, kv_connector_extra_config: %r," - " server block size: %d, vllm block size: %d,", - self._extra_config.__dict__, - self._manager_block_size, - self._local_block_size, - ) + "TairKvCacheConnector scheduler inited, extra_config: %r, manager block size: %d, " + "vllm block size: %d, groups: %d", + self._extra_config.__dict__, self._manager_block_size, + self._vllm_block_size, self._num_groups) elif role == KVConnectorRole.WORKER: self._tp_rank = get_tensor_model_parallel_rank() self._device_mod = None - self._save_stream = None - self._load_stream = None - - logger.warning( - "TairKvCacheConnector in worker inited, tp rank: %d, tp size: %d, host_ip: %s, port: %d" % ( - self._tp_rank, self._tp_size, self._host_ip, port) - ) - if self._tp_rank == 0: - # start coordinator - self._coordinator_server = TpCoordinatorServer(self._host_ip, port, self._tp_size, - self.on_save_finished) - + self._coordinator_server = TpCoordinatorServer( + self._host_ip, port, self._tp_size, self.on_save_finished) self._coordinator_client = TpCoordinatorClient(self._host_ip, port) self._storage_configs = register_response["storage_configs"] - # data transfer setup - self._location_spec_name = self._tp_rank_to_spec_name(self._tp_rank) - self._write_timeout_seconds = self._extra_config.write_timeout_seconds - - sdk_backend_configs = [] - - hf3fs_configs = self.parse_hf3fs_configs(self._storage_configs) - sdk_backend_configs.extend(hf3fs_configs) - logger.debug(sdk_backend_configs) + sdk_backend_configs = self.parse_hf3fs_configs(self._storage_configs) + self._self_spec_names = { + meta.group_idx: self._spec_name(self._tp_rank, meta.group_idx) + for meta in self._group_metas + } transfer_client_json = { "instance_group": self._extra_config.instance_group, "instance_id": self._extra_config.instance_id, @@ -257,26 +270,63 @@ def __init__(self, }, }, "location_spec_infos": { - self._location_spec_name: math.prod(per_manager_location_spec_shape) * kv_dtype.itemsize, + self._self_spec_names[meta.group_idx]: meta.per_block_bytes + for meta in self._group_metas }, } - self._transfer_client_config = json.dumps(transfer_client_json) - - self._init_params = kvcm_py_client.InitParams() - self._init_params.role_type = kvcm_py_client.RoleType.WORKER - self._init_params.self_location_spec_name = self._location_spec_name - self._init_params.storage_configs = f"{self._storage_configs}" - - logger.info("_transfer_client_config:%s, _init_params:%s", self._transfer_client_config, self._init_params) - + init_params = kvcm_py_client.InitParams() + init_params.role_type = kvcm_py_client.RoleType.WORKER + init_params.self_location_spec_name = self._self_spec_names[self._group_metas[0].group_idx] + init_params.storage_configs = f"{self._storage_configs}" + transfer_client_config = json.dumps(transfer_client_json) + logger.info("transfer_client_config: %s", transfer_client_config) self._transfer_client = kvcm_py_client.TransferClient.Create( - self._transfer_client_config, self._init_params - ) + transfer_client_config, init_params) assert self._transfer_client is not None, "kvcm_py_client.TransferClient.Create failed" + logger.warning( + "TairKvCacheConnector worker inited, tp rank: %d/%d, host: %s:%d, groups: %d", + self._tp_rank, self._tp_size, self._host_ip, port, self._num_groups) + + def _spec_name(self, tp_rank: int, group_idx: int) -> str: + return f"tp{tp_rank}_g{group_idx}" + + def _parse_groups(self, kv_cache_config: "KVCacheConfig") -> List[GroupMeta]: + metas = [] + for idx, group in enumerate(kv_cache_config.kv_cache_groups): + if getattr(group, "is_eagle_group", False): + logger.warning("skip eagle group %d (%d layers)", idx, len(group.layer_names)) + continue + spec = group.kv_cache_spec + if isinstance(spec, MambaSpec): + metas.append(GroupMeta( + group_idx=idx, + is_attention=False, + layer_names=list(group.layer_names), + block_size=spec.block_size, + per_block_bytes=spec.page_size_bytes * len(group.layer_names), + page_size_bytes=spec.page_size_bytes, + )) + elif isinstance(spec, FullAttentionSpec): + # Attention KV is token-granular; scale from the spec's page size + # to the manager block size. + per_token_bytes = spec.page_size_bytes // spec.block_size + metas.append(GroupMeta( + group_idx=idx, + is_attention=True, + layer_names=list(group.layer_names), + block_size=spec.block_size, + per_block_bytes=per_token_bytes * self._manager_block_size * len(group.layer_names), + )) + else: + raise NotImplementedError( + f"Unsupported kv cache spec {type(spec).__name__} in group {idx}") + assert metas, "no usable kv cache groups" + return metas def shutdown(self): - # TODO: stop background threads and cleanup transfer client self._manager_client.close() + if hasattr(self, "_location_query_manager"): + self._location_query_manager.shutdown() return None def parse_hf3fs_configs(self, storage_configs): @@ -286,7 +336,7 @@ def parse_hf3fs_configs(self, storage_configs): if storage_config["type"] == "vcns_hf3fs": storage_config["type"] = "hf3fs" if storage_config["type"] == "hf3fs" and storage_config["is_available"]: - hf3fs_config = { + hf3fs_configs.append({ "type": storage_config["type"], "mountpoint": storage_config["storage_spec"]["mountpoint"], "root_dir": storage_config["storage_spec"]["root_dir"], @@ -294,371 +344,398 @@ def parse_hf3fs_configs(self, storage_configs): "read_iov_size": self._iov_size, "write_iov_block_size": self._extra_config.write_iov_block_size, "write_iov_size": self._iov_size, - } - hf3fs_configs.append(hf3fs_config) + }) self._storage_configs = json.dumps(storage_configs_json) return hf3fs_configs - def generate_blocks(self, token_ids, block_size, max_token_length) -> list[dict[str, Any]]: - results = [] - token_length = min(len(token_ids), max_token_length) - # token_length = len(token_ids) - for i in range(0, token_length, block_size): - if i + block_size > token_length: - break - results.append({ - "token_ids": token_ids[i:i + block_size], - "unique_id": None, - "location": None - }) - return results - - # ============================== - # Worker-side methods - # ============================== - - def generate_blocks_idx(self, manager_block_idxes, local_block_ids): - blocks_idx = [] - for manager_block_idx in manager_block_idxes: - # get kvcache index list - block_idx = [] - for i in range(self._manager_block_size): - now_token_idx = manager_block_idx * self._manager_block_size + i - assert now_token_idx // self._local_block_size < len(local_block_ids) - local_block_id = local_block_ids[now_token_idx // self._local_block_size] - token_offset = now_token_idx % self._local_block_size - block_idx.append(local_block_id * self._local_block_size + token_offset) - blocks_idx.append(block_idx) - return blocks_idx - - def on_save_finished(self, write_session_id: str, save_context: SaveContext): - logger.debug(save_context.result_per_rank) - for block_idx in range(len(save_context.locations)): - # TODO: report uri when enable local alloc - # location_specs = [] - is_fully_saved = True - for rank in range(self._tp_size): - is_success = save_context.result_per_rank[rank][block_idx] - if not is_success: - # this spec is not fully saved, report failed - is_fully_saved = False - # else: - # # Convert the spec to include name field instead of tp_rank - # location_specs.append({ - # "name": self._tp_rank_to_spec_name(rank), - # "uri": spec - # }) - if is_fully_saved: - # save_context.locations[block_idx]["location_specs"] = location_specs - save_context.success_mask.append(True) - else: - save_context.success_mask.append(False) - logger.debug("finish_write_cache blocks:%s mask:%s write_session_id:%s", save_context.locations, - save_context.success_mask, write_session_id) - try: - self._manager_client.finish_write_cache({ - "trace_id": "test_test", - "instance_id": self._extra_config.instance_id, - "write_session_id": write_session_id, - "success_blocks": { - "bool_masks": { - "values": save_context.success_mask - } - } - }) - except Exception as e: - logger.warning("finish_write_cache failed, write_session_id: %s, error: %s", write_session_id, e) - + # ------------------------------------------------------------------ # + # Worker side: KV cache registration + # ------------------------------------------------------------------ # def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): - _, first_layer_kvcache = next(iter(kv_caches.items())) self._kv_caches = kv_caches - # TODO: support MLA - - assert self._local_block_size == first_layer_kvcache.shape[2], "kv cache shape error" - for layer_name, kvcache in kv_caches.items(): - assert kvcache.is_contiguous(), "kv cache must be contiguous" - - # torch.Size([2, block_num, block_size, kv_head_num, kv_dim]) - # 2 -> key, value - self._local_block_num = first_layer_kvcache.shape[1] - self._local_token_num = self._local_block_num * self._local_block_size - - self._dtype = first_layer_kvcache.dtype - self._device = first_layer_kvcache.device + first_attn = next(kv_caches[name] + for meta in self._group_metas if meta.is_attention + for name in meta.layer_names) + self._dtype = first_attn.dtype + self._device = first_attn.device self._device_mod = _get_device_module(self._device) - self._save_stream = self._device_mod.Stream() - self._load_stream = self._device_mod.Stream() - self._per_manager_location_spec_layer_shape = [first_layer_kvcache.shape[0], - self._manager_block_size, - first_layer_kvcache.shape[3] * first_layer_kvcache.shape[4]] - self._per_manager_location_spec_layer_byte_size = math.prod( - self._per_manager_location_spec_layer_shape) * self._dtype.itemsize - self._per_layer_token_key_dim_size = first_layer_kvcache.shape[3] * first_layer_kvcache.shape[4] - self._per_layer_token_key_byte_size = (first_layer_kvcache.shape[3] * - first_layer_kvcache.shape[4] * self._dtype.itemsize) - assert self._per_layer_token_key_byte_size == first_layer_kvcache[0][0][1].data_ptr() - \ - first_layer_kvcache[0][0][0].data_ptr(), "kv cache shape error" - assert self._per_manager_location_spec_layer_byte_size == 2 * self._manager_block_size * self._per_layer_token_key_byte_size - - self._per_manager_location_spec_shape = [len(self._kv_caches)] + self._per_manager_location_spec_layer_shape - self._per_manager_location_spec_byte_size = math.prod( - self._per_manager_location_spec_shape) * self._dtype.itemsize - - self._kvcache_ptr_tensor_cpu = torch.tensor( - [self._kv_caches[name].data_ptr() for name in self._kv_caches], - dtype=torch.int64, - device="cpu" - ) - self._kvcache_ptr_tensor_gpu = self._kvcache_ptr_tensor_cpu.to(self._device) - if self._use_mla: - self._all_kvcache_ptr_tensor_cpu = torch.tensor( - [self._kv_caches[name].data_ptr() for name in self._kv_caches], - dtype=torch.int64, - device="cpu" - ) - else: - kvcache_ptrs = [] - for name in self._kv_caches: - kvcache_ptrs.append(self._kv_caches[name][0].data_ptr()) - kvcache_ptrs.append(self._kv_caches[name][1].data_ptr()) - self._all_kvcache_ptr_tensor_cpu = torch.tensor( - kvcache_ptrs, - dtype=torch.int64, - device="cpu" - ) - self._all_kvcache_ptr_tensor_gpu = self._all_kvcache_ptr_tensor_cpu.to(self._device) + + groups = [self._build_transfer_group(meta, kv_caches) for meta in self._group_metas] self._kvcache_info = KVCacheInfo( - self._tp_rank, - self._tp_size, - self._kv_caches, - self._kvcache_ptr_tensor_cpu, - self._kvcache_ptr_tensor_gpu, - self._all_kvcache_ptr_tensor_gpu, - len(self._kv_caches), - self._local_token_num, - tuple(self._per_manager_location_spec_shape), - self._per_manager_location_spec_byte_size, - self._per_layer_token_key_dim_size, - self._device, - self._dtype + tp_rank=self._tp_rank, + world_size=self._tp_size, + groups=groups, + device=self._device, + dtype=self._dtype, ) - self._copy_buffer_allocator = CopyBufferAllocator(torch.device("cpu"), self._dtype, - self._per_manager_location_spec_shape, 1024) - - # 初始化DataTransferManager实例 self._data_transfer = DataTransferManager( - self._kvcache_info, - self._manager_block_size, - self._copy_buffer_allocator, - self._transfer_client, - self._coordinator_client, - self._extra_config, - ) + self._kvcache_info, self._manager_block_size, + self._transfer_client, self._coordinator_client, self._extra_config) + + logger.warning("register_kv_caches done: %s", [ + (g.spec_name, "attn" if g.is_attention else "state", + g.layer_num, g.per_block_bytes) for g in groups]) + + def _build_transfer_group(self, meta: GroupMeta, kv_caches) -> TransferGroup: + spec_name = self._self_spec_names[meta.group_idx] + if meta.is_attention: + tensors = [kv_caches[name] for name in meta.layer_names] + ref = tensors[0] + # vLLM >= 0.26.0 packs K and V into the content dim: logical shape + # (num_blocks, num_kv_heads, kernel_block_size, 2*head_size). With the + # default NHD stride order the memory is laid out token-major as + # (num_blocks, kernel_block_size, num_kv_heads, 2*head_size), so a flat + # index (global_token * per_token_dim + dim) walks the storage + # correctly. K/V packing is opaque to the byte-exact transport. + assert ref.dim() == 4, f"unexpected kv layout {ref.shape}" + for t in tensors: + assert t.shape == ref.shape and t.stride() == ref.stride(), \ + "attention layers in one group must share shape/stride" + kernel_block_size = ref.shape[2] + assert meta.block_size % kernel_block_size == 0, \ + f"group block size {meta.block_size} not a multiple of kernel " \ + f"block size {kernel_block_size}" + per_token_dim = ref.shape[1] * ref.shape[3] # num_kv_heads * 2*head_size + # The gather/scatter kernel needs token-major memory inside a page: + # dims (blk, head, tok, dim) laid out as (blk, tok, head, dim). This + # is vLLM's NHD order; HND would interleave heads across tokens. + assert ref.stride()[1:] == (ref.shape[3], per_token_dim, 1), \ + f"kv cache page not token-major: shape={ref.shape} " \ + f"stride={ref.stride()}; set VLLM_KV_CACHE_LAYOUT=NHD" + # Padded pages (page_size_padded) leave gaps between blocks; the + # kernel's strided path skips them. Stride 0 = fast flat indexing. + flat = ref.stride(0) == kernel_block_size * per_token_dim + block_stride = 0 if flat else ref.stride(0) + # One pointer per layer: K and V are packed in the content dim, and + # data_ptr() of the permuted view is the storage base. + ptrs = [t.data_ptr() for t in tensors] + ptr_tensor = torch.tensor(ptrs, dtype=torch.int64, device="cpu").to(self._device) + return TransferGroup( + group_idx=meta.group_idx, + spec_name=spec_name, + is_attention=True, + layer_names=meta.layer_names, + block_size=meta.block_size, + per_block_bytes=meta.per_block_bytes, + kvcache_ptr_tensor_gpu=ptr_tensor, + layer_num=len(meta.layer_names), + per_token_dim=per_token_dim, + kernel_block_size=kernel_block_size, + kv_stride=0, + block_stride=block_stride, + ) - logger.warning("register_kv_caches, _per_manager_location_spec_layer_shape: %s", - self._per_manager_location_spec_layer_shape) + # Mamba/state group: each layer is a list[Tensor] sharing one storage; + # rebuild a (num_blocks, page_size_bytes) byte view for opaque copy. + block_views = [] + for name in meta.layer_names: + states = kv_caches[name] + assert isinstance(states, (list, tuple)) and len(states) > 0, \ + f"state layer {name} should be a list of tensors" + storage = states[0].untyped_storage() + for st in states[1:]: + assert st.untyped_storage().data_ptr() == storage.data_ptr(), \ + f"state layer {name}: tensors do not share storage" + num_blocks = states[0].shape[0] + need = num_blocks * meta.page_size_bytes + assert storage.nbytes() >= need, \ + f"state layer {name}: storage {storage.nbytes()} < {need}" + byte_view = torch.tensor([], dtype=torch.uint8, device=self._device).set_(storage) + block_views.append(byte_view[:need].view(num_blocks, meta.page_size_bytes)) + return TransferGroup( + group_idx=meta.group_idx, + spec_name=spec_name, + is_attention=False, + layer_names=meta.layer_names, + block_size=meta.block_size, + per_block_bytes=meta.per_block_bytes, + layer_num=len(meta.layer_names), + block_view_tensors=block_views, + page_size_bytes=meta.page_size_bytes, + ) + # ------------------------------------------------------------------ # + # Block index translation + # ------------------------------------------------------------------ # + def _attn_token_indices(self, group: TransferGroup, manager_block_idxes, + block_table) -> List[List[int]]: + """Map manager blocks to flat token slots of one attention group. + + Three-tier hierarchy: + manager block (KVCM unit) -> global token idx + -> group block (block_table unit, group.block_size tokens) + -> kernel physical block (tensor unit; ratio physical per group block). + """ + mbs = self._manager_block_size + gbs = group.block_size + kbs = group.kernel_block_size + ratio = gbs // kbs + out = [] + for mb in manager_block_idxes: + idxs = [] + base = mb * mbs + for i in range(mbs): + tok = base + i + logical = tok // gbs + assert logical < len(block_table), ( + f"group block {logical} out of range (len={len(block_table)})") + off = tok % gbs + phys = block_table[logical] * ratio + off // kbs + idxs.append(phys * kbs + off % kbs) + out.append(idxs) + return out + + def _state_block_ids(self, group: TransferGroup, manager_block_idxes, + block_table) -> List[int]: + """Map manager blocks to block ids of a state (mamba) group. + + State is stored once per group block and covers the whole prefix up to + that block, so the manager block's last token selects the block.""" + mbs = self._manager_block_size + gbs = group.block_size + out = [] + for mb in manager_block_idxes: + logical = ((mb + 1) * mbs - 1) // gbs + assert logical < len(block_table), ( + f"group block {logical} out of range (len={len(block_table)})") + out.append(block_table[logical]) + return out + + def _self_uris(self, locations, spec_name: str) -> List[str]: + uris = [] + for location in locations: + for spec in location.get("location_specs", []): + if spec["name"] == spec_name: + uris.append(spec["uri"]) + return uris + + # ------------------------------------------------------------------ # + # Worker side: load / save + # ------------------------------------------------------------------ # + def _submit_group_tasks(self, task_fn, multi_result, task_idx, group, + uris, token_indices, block_ids, per_task_size, *extra): + for i in range(0, len(uris), per_task_size): + end = min(len(uris), i + per_task_size) + self._data_transfer.submit_task( + task_fn, multi_result, task_idx, group, uris[i:end], + token_indices[i:end] if token_indices is not None else None, + block_ids[i:end] if block_ids is not None else None, *extra) + task_idx += 1 + return task_idx + + def _plan_group_transfers(self, locations, manager_block_idxes, block_ids_per_group): + """Build (group, uris, token_indices, block_ids) for every group. + + Returns None if any group's URI list does not cover all blocks.""" + num_blocks = len(manager_block_idxes) + plans = [] + for group in self._kvcache_info.groups: + uris = self._self_uris(locations, group.spec_name) + if len(uris) != num_blocks: + logger.warning("group %s: %d uris for %d blocks, skip transfer", + group.spec_name, len(uris), num_blocks) + return None + # block_ids_per_group is indexed by the vLLM group index. + block_table = block_ids_per_group[group.group_idx] + if group.is_attention: + plans.append((group, uris, + self._attn_token_indices(group, manager_block_idxes, block_table), + None)) + else: + plans.append((group, uris, None, + self._state_block_ids(group, manager_block_idxes, block_table))) + return plans def start_load_kv(self, forward_context: "ForwardContext", **kwargs) -> None: meta = typing.cast(TairKvCacheConnectorMetadata, self._get_connector_metadata()) - for load_req in meta.to_load_requests: - if len(load_req.need_load_locations) == 0: + if not load_req.need_load_locations: + continue + num_blocks = len(load_req.manager_block_idxes) + plans = self._plan_group_transfers( + load_req.need_load_locations, load_req.manager_block_idxes, + load_req.all_block_ids) + + # Report failures against the block table vLLM can act on: map each + # manager block to the logical block holding its first token. vLLM + # truncates computed tokens at the first invalid block, so this is + # sufficient for recovery. vLLM's invalid-block handling only + # supports single-group models; for hybrid models a failed load can + # only be logged. + report_ids = [] + if self._num_groups == 1: + table = load_req.all_block_ids[0] + gbs = self._group_metas[0].block_size + report_ids = [table[(mb * self._manager_block_size) // gbs] + for mb in load_req.manager_block_idxes] + done_cb = self._data_transfer.create_load_done_callback( + load_req.req_id, self._tp_rank, meta.epoch, + copy.copy(report_ids), num_blocks, + report_failures=self._num_groups == 1) + + if plans is None: + # Nothing submitted; report the whole load as failed. + mr = MultiResult(1, done_cb) + mr.submit_result(0, [False] * num_blocks * self._num_groups) continue - block_token_indices = self.generate_blocks_idx(load_req.manager_block_idxes, load_req.local_block_ids) - all_remote_uris = self.get_self_uris(load_req.need_load_locations) - - per_task_size = self._extra_config.block_per_load_task - task_num = math.ceil(len(block_token_indices) / per_task_size) - done_callback = self._data_transfer.create_load_done_callback( - load_req.req_id, - self._kvcache_info.tp_rank, - meta.epoch, - copy.copy(load_req.local_block_ids) - ) - multi_result = MultiResult(task_num, done_callback) - + per_task = self._extra_config.block_per_load_task + task_num = sum(math.ceil(num_blocks / per_task) for _ in plans) + multi_result = MultiResult(task_num, done_cb) task_idx = 0 - for i in range(0, len(block_token_indices), per_task_size): - end_idx = min(len(block_token_indices), i + per_task_size) - task_remote_uris = all_remote_uris[i:end_idx] - task_block_token_indices = block_token_indices[i:end_idx] - self._data_transfer.submit_task(self._data_transfer.load_task, multi_result, task_idx, task_remote_uris, - task_block_token_indices) - task_idx += 1 + for group, uris, token_indices, block_ids in plans: + task_idx = self._submit_group_tasks( + self._data_transfer.load_task, multi_result, task_idx, + group, uris, token_indices, block_ids, per_task) def wait_for_layer_load(self, layer_name: str) -> None: - # logger.warning("wait_for_layer_load, layer_name: %s", layer_name) pass - def save_kv_layer(self, layer_name: str, kv_layer: torch.Tensor, attn_metadata: "AttentionMetadata", - **kwargs) -> None: - # logger.warning("save_kv_layer, layer_name: %s", layer_name) + def save_kv_layer(self, layer_name: str, kv_layer: torch.Tensor, + attn_metadata: "AttentionMetadata", **kwargs) -> None: pass def wait_for_save(self): meta = typing.cast(TairKvCacheConnectorMetadata, self._get_connector_metadata()) - # logger.warning("wait_for_save, meta: %r", meta) - - kvcache_ready_event = None - if len(meta.to_save_requests) > 0: - kvcache_ready_event = self._device_mod.Event() - kvcache_ready_event.record(self._device_mod.current_stream()) - - for req_save in meta.to_save_requests: - req = self._alive_requests[req_save.req_id] - - # get idx - blocks_idx = self.generate_blocks_idx(req_save.manager_block_idxes, req.local_block_ids) - all_remote_uris = self.get_self_uris(req_save.target_locations) - - per_task_size = self._extra_config.block_per_save_task - task_num = math.ceil(len(blocks_idx) / per_task_size) - done_callback = self._data_transfer.create_save_done_callback( - req.req_id, - self._kvcache_info.tp_rank, - req_save.write_session_id - ) - multi_result = MultiResult(task_num, done_callback) + if not meta.to_save_requests: + return + ready_event = self._device_mod.Event() + ready_event.record(self._device_mod.current_stream()) + + for save_req in meta.to_save_requests: + req = self._alive_requests[save_req.req_id] + num_blocks = len(save_req.manager_block_idxes) + plans = self._plan_group_transfers( + save_req.target_locations, save_req.manager_block_idxes, + req.block_ids_per_group) + + done_cb = self._data_transfer.create_save_done_callback( + req.req_id, self._tp_rank, save_req.write_session_id, num_blocks) + + if plans is None: + mr = MultiResult(1, done_cb) + mr.submit_result(0, [False] * num_blocks * self._num_groups) + continue + per_task = self._extra_config.block_per_save_task + task_num = sum(math.ceil(num_blocks / per_task) for _ in plans) + multi_result = MultiResult(task_num, done_cb) task_idx = 0 - for i in range(0, len(blocks_idx), per_task_size): - end_idx = min(len(blocks_idx), i + per_task_size) - task_remote_uris = all_remote_uris[i:end_idx] - task_block_token_indices = blocks_idx[i:end_idx] - self._data_transfer.submit_task(self._data_transfer.save_task, multi_result, task_idx, task_remote_uris, - task_block_token_indices, - kvcache_ready_event) - task_idx += 1 + for group, uris, token_indices, block_ids in plans: + task_idx = self._submit_group_tasks( + self._data_transfer.save_task, multi_result, task_idx, + group, uris, token_indices, block_ids, per_task, ready_event) if self._tp_rank == 0: req.scheduled_saving_count += 1 - def get_self_uris(self, locations): - all_remote_uris = [] - for idx, location in enumerate(locations): - for location_spec in location["location_specs"]: - # Match by location spec name instead of tp_rank - if self._tp_rank_to_spec_name(self._kvcache_info.tp_rank) == location_spec["name"]: - all_remote_uris.append(location_spec["uri"]) - return all_remote_uris - - def get_finished( - self, finished_req_ids: set[str] - ) -> Tuple[Optional[set[str]], Optional[set[str]]]: + def on_save_finished(self, write_session_id: str, save_context: SaveContext): + for block_idx in range(len(save_context.locations)): + fully_saved = all(save_context.result_per_rank[rank][block_idx] + for rank in range(self._tp_size)) + save_context.success_mask.append(fully_saved) + logger.debug("finish_write_cache mask:%s session:%s", + save_context.success_mask, write_session_id) + try: + self._manager_client.finish_write_cache({ + "trace_id": "finish_%s" % write_session_id[:8], + "instance_id": self._extra_config.instance_id, + "write_session_id": write_session_id, + "success_blocks": {"bool_masks": {"values": save_context.success_mask}}, + }) + except Exception as e: + logger.warning("finish_write_cache failed, session: %s, error: %s", + write_session_id, e) + + def get_finished(self, finished_req_ids: set) -> Tuple[Optional[set], Optional[set]]: meta = typing.cast(TairKvCacheConnectorMetadata, self._get_connector_metadata()) if self._tp_rank != 0: for finish_req in meta.to_finish_requests: - req_id = finish_req.req_id - if req_id in self._alive_requests: - self._alive_requests.pop(req_id) + self._alive_requests.pop(finish_req.req_id, None) return None, None - # self._tp_rank == 0 - finished_saving_reqs = [] - # check if any request is saving kvcache - (finished_saving_tasks, finished_loading_tasks) = self._coordinator_server.get_finished_tasks() + finished_saving = [] + finished_saving_tasks, finished_loading_tasks = self._coordinator_server.get_finished_tasks() for req_id in finished_saving_tasks: req = self._alive_requests[req_id] req.sent_saving_count += 1 - assert req.sent_saving_count <= req.scheduled_saving_count if (req.need_report_after_saving_finished and req.sent_saving_count == req.scheduled_saving_count): - finished_saving_reqs.append(req_id) + finished_saving.append(req_id) self._alive_requests.pop(req_id) for finish_req in meta.to_finish_requests: - req_id = finish_req.req_id - if req_id not in self._alive_requests: - # called get_num_new_matched_tokens but never scheduled + req = self._alive_requests.get(finish_req.req_id) + if req is None: continue - req = self._alive_requests[req_id] if req.sent_saving_count == req.scheduled_saving_count: - finished_saving_reqs.append(req_id) - self._alive_requests.pop(req_id) + finished_saving.append(req.req_id) + self._alive_requests.pop(req.req_id) else: - self._alive_requests[req_id].need_report_after_saving_finished = True - return set(finished_saving_reqs), set(finished_loading_tasks) + req.need_report_after_saving_finished = True + return set(finished_saving), set(finished_loading_tasks) - def get_block_ids_with_load_errors(self) -> set[int]: + def get_block_ids_with_load_errors(self) -> set: if self._tp_rank != 0: return set() - failed_set = self._coordinator_server.get_failed_loading_block_idxs() - if len(failed_set) > 0: - logger.warning("block_ids_with_load_errors: %s", failed_set) - return failed_set + failed = self._coordinator_server.get_failed_loading_block_idxs() + if failed: + logger.warning("block_ids_with_load_errors: %s", failed) + return failed - def bind_connector_metadata( - self, connector_metadata: KVConnectorMetadata) -> None: + def bind_connector_metadata(self, connector_metadata: KVConnectorMetadata) -> None: self._connector_metadata = connector_metadata meta = typing.cast(TairKvCacheConnectorMetadata, self._get_connector_metadata()) - - for req_state_delta in meta.requests: - if req_state_delta.req_id not in self._alive_requests: - assert not req_state_delta.is_delta - if not req_state_delta.is_delta: - self._alive_requests[req_state_delta.req_id] = ReqState.create_from_delta(req_state_delta) + for delta in meta.requests: + if not delta.is_delta: + self._alive_requests[delta.req_id] = ReqState.create_from_delta(delta) else: - self._alive_requests[req_state_delta.req_id].update_from_delta(req_state_delta) - - # ============================== - # Scheduler-side methods - # ============================== - def get_num_new_matched_tokens(self, request: "Request", num_computed_tokens: int) -> Tuple[int, bool]: - # logger.warning("get matched token ids: %s, id: %s", request.prompt_token_ids, request.request_id) - - bypass_match = False - # TODO: add arrival_time to req_id in order to handle same request id - if request.request_id in self._alive_requests: - # bypass remote match for alive requests - # possible cases: - # 1. reschedule when all kvcache loading failed - # 2. TODO: no enough hbm to schedule the request - # bypass_match = True - # logger.warning("bypass match for alive request, req_id: %s", request.request_id) - pass - - computed_manager_block_size = num_computed_tokens // self._manager_block_size - all_calced_remote_block_num = computed_manager_block_size - new_matched_count = 0 - - if not bypass_match: - is_query_done, need_load_locations = ( - self._location_query_manager.get_locations_for_query(request, computed_manager_block_size)) - if not is_query_done: - # async get_cache_location - return None, False - new_matched_count = len(need_load_locations) * self._manager_block_size - logger.info("req:%s, new_matched_count:%d", request.request_id, new_matched_count) - - all_calced_remote_block_num = computed_manager_block_size + len(need_load_locations) - - if new_matched_count != 0: - self._waiting_to_load_requests.append(LoadRequest( - req_id=request.request_id, - manager_block_idxes=[i for i in range(computed_manager_block_size, all_calced_remote_block_num)], - need_load_locations=need_load_locations, - )) - - new_req_meta = ReqState(request.request_id, copy.copy(request.prompt_token_ids), [], - all_calced_remote_block_num, - num_computed_tokens, - new_matched_count, - request) + assert delta.req_id in self._alive_requests + self._alive_requests[delta.req_id].update_from_delta(delta) + + # ------------------------------------------------------------------ # + # Scheduler side + # ------------------------------------------------------------------ # + def get_num_new_matched_tokens(self, request: "Request", + num_computed_tokens: int) -> Tuple[Optional[int], bool]: + computed_blocks = num_computed_tokens // self._manager_block_size + + is_query_done, need_load_locations = ( + self._location_query_manager.get_locations_for_query(request, computed_blocks)) + if not is_query_done: + # async query in flight; vLLM will ask again + return None, False + + new_matched_count = len(need_load_locations) * self._manager_block_size + total_remote_blocks = computed_blocks + len(need_load_locations) + logger.info("req:%s matched %d external tokens", request.request_id, new_matched_count) + + if new_matched_count: + self._waiting_to_load_requests.append(LoadRequest( + req_id=request.request_id, + manager_block_idxes=list(range(computed_blocks, total_remote_blocks)), + need_load_locations=need_load_locations, + )) - self._alive_requests[request.request_id] = new_req_meta + self._alive_requests[request.request_id] = ReqState( + req_id=request.request_id, + token_ids=copy.copy(request.prompt_token_ids), + block_ids_per_group=[], + has_saved_block_num=total_remote_blocks, + local_matched_token_num=num_computed_tokens, + remote_matched_token_num=new_matched_count, + vllm_request=request, + ) return new_matched_count, new_matched_count > 0 - def update_state_after_alloc(self, request: "Request", blocks: "KVCacheBlocks", num_external_tokens: int): - if request.request_id not in self._alive_requests: + def update_state_after_alloc(self, request: "Request", blocks: "KVCacheBlocks", + num_external_tokens: int): + req_state = self._alive_requests.get(request.request_id) + if req_state is None: return - req_state = self._alive_requests[request.request_id] - # blocks_ids[0]: only one KV cache groups for now - # refer to vllm/v1/core/kv_cache_manager.py:35 - req_state.local_block_ids = copy.copy(blocks.get_block_ids()[0]) + req_state.block_ids_per_group = [list(b) for b in blocks.get_block_ids()] def build_connector_meta(self, scheduler_output: SchedulerOutput) -> KVConnectorMetadata: meta = TairKvCacheConnectorMetadata(self._epoch) @@ -666,89 +743,78 @@ def build_connector_meta(self, scheduler_output: SchedulerOutput) -> KVConnector for load_req in self._waiting_to_load_requests: request = self._alive_requests[load_req.req_id] - if len(request.local_block_ids) == 0: - # ignore load_req if vllm has not called update_state_after_alloc, - # vllm will call get_num_new_matched_tokens again + if not request.block_ids_per_group: + # update_state_after_alloc was never called; vLLM will re-query. continue - load_req.local_block_ids = request.local_block_ids + load_req.all_block_ids = [list(b) for b in request.block_ids_per_group] meta.add_load_request(load_req) self._waiting_to_load_requests = [] for vllm_req in scheduler_output.scheduled_new_reqs: request = self._alive_requests[vllm_req.req_id] - request.local_block_ids = copy.copy(vllm_req.block_ids[0]) - - state_to_worker = ReqStateToWorker(req_id=request.req_id, - has_saved_block_num=request.has_saved_block_num, - new_tokens_ids=request.token_ids, - new_local_block_ids=request.local_block_ids, - is_delta=False - ) - meta.add_req_state_to_worker(state_to_worker) - logger.info("new request: %s, block_ids_len: %d", vllm_req.req_id, len(vllm_req.block_ids[0])) + request.block_ids_per_group = [list(b) for b in vllm_req.block_ids] + meta.add_req_state_to_worker(ReqStateToWorker( + req_id=request.req_id, + has_saved_block_num=request.has_saved_block_num, + new_tokens_ids=request.token_ids, + new_block_ids_per_group=request.block_ids_per_group, + is_delta=False, + )) cached_reqs = scheduler_output.scheduled_cached_reqs for idx, req_id in enumerate(cached_reqs.req_ids): request = self._alive_requests[req_id] - vllm_req = request.vllm_request num_new_tokens = scheduler_output.num_scheduled_tokens[req_id] num_current_tokens = len(request.token_ids) - - new_token_ids = vllm_req.all_token_ids[ - num_current_tokens: num_current_tokens + num_new_tokens - ] - state_to_worker = ReqStateToWorker(req_id=request.req_id, - has_saved_block_num=request.has_saved_block_num) - + new_token_ids = request.vllm_request.all_token_ids[ + num_current_tokens:num_current_tokens + num_new_tokens] request.token_ids.extend(new_token_ids) - state_to_worker.new_tokens_ids = new_token_ids - resumed_from_preemption = False - if hasattr(cached_reqs, "resumed_req_ids"): - # vllm >= 0.11.1 - resumed_from_preemption = req_id in cached_reqs.resumed_req_ids - else: - # vllm <= 0.11.0 - resumed_from_preemption = cached_reqs.resumed_from_preemption[idx] + delta = ReqStateToWorker( + req_id=request.req_id, + has_saved_block_num=request.has_saved_block_num, + new_tokens_ids=new_token_ids, + ) - if resumed_from_preemption: - request.local_block_ids = copy.copy(cached_reqs.new_block_ids[idx][0]) - state_to_worker.resumed_from_preemption = True - state_to_worker.new_local_block_ids = request.local_block_ids + if hasattr(cached_reqs, "resumed_req_ids"): + resumed = req_id in cached_reqs.resumed_req_ids else: - if cached_reqs.new_block_ids[idx] is None: - # https://github.com/vllm-project/vllm/pull/23262 - continue - new_block_ids = cached_reqs.new_block_ids[idx][0] - request.local_block_ids.extend(new_block_ids) - state_to_worker.new_local_block_ids = new_block_ids - meta.add_req_state_to_worker(state_to_worker) + resumed = cached_reqs.resumed_from_preemption[idx] + + new_block_ids = cached_reqs.new_block_ids[idx] + if resumed: + request.block_ids_per_group = [list(b) for b in new_block_ids] + delta.resumed_from_preemption = True + delta.new_block_ids_per_group = request.block_ids_per_group + elif new_block_ids is not None: + # https://github.com/vllm-project/vllm/pull/23262: may be None + delta.new_block_ids_per_group = [list(b) for b in new_block_ids] + for group_ids, new_ids in zip(request.block_ids_per_group, + delta.new_block_ids_per_group): + group_ids.extend(new_ids) + meta.add_req_state_to_worker(delta) for req in self._alive_requests.values(): - target_save_num = min(len(req.token_ids), - len(req.local_block_ids) * self._local_block_size) // self._manager_block_size + target_save_num = min( + len(req.token_ids), + req.num_allocated_blocks * self._vllm_block_size) // self._manager_block_size if target_save_num > req.has_saved_block_num: req.scheduled_saving_count += 1 self._http_executor.submit( - self.start_save_kvcache_async, - req.req_id, + self.start_save_kvcache_async, req.req_id, req.token_ids[:target_save_num * self._manager_block_size], - target_save_num - ) + target_save_num) req.has_saved_block_num = target_save_num - new_save_reqs: List[SaveRequest] = [] with self._waiting_to_save_requests_lock: new_save_reqs = self._waiting_to_save_requests self._waiting_to_save_requests = [] for save_req in new_save_reqs: - if save_req.req_id not in self._alive_requests: - # TODO: should not happen anymore + req = self._alive_requests.get(save_req.req_id) + if req is None: logger.warning("request %s is not alive, skip saving", save_req.req_id) continue - req = self._alive_requests[save_req.req_id] meta.add_save_request(save_req) - req.sent_saving_count += 1 if (req.need_report_after_saving_finished and req.scheduled_saving_count == req.sent_saving_count): @@ -760,8 +826,6 @@ def build_connector_meta(self, scheduler_output: SchedulerOutput) -> KVConnector for finish_req in self._waiting_to_finish_requests: meta.add_finish_request(finish_req) self._waiting_to_finish_requests = [] - - # logger.warning("build_connector_meta: %r", meta) return meta def start_save_kvcache_async(self, req_id, token_ids, target_save_num): @@ -770,62 +834,49 @@ def start_save_kvcache_async(self, req_id, token_ids, target_save_num): "instance_id": self._extra_config.instance_id, "block_keys": [], "token_ids": token_ids, - "write_timeout_seconds": 30 + "write_timeout_seconds": self._extra_config.write_timeout_seconds, } - logger.debug("start_write_cache req: %s", request) try: response = self._manager_client.start_write_cache(request) except Exception as e: - logger.warning("start_write_cache error, skip this saving, exception: %s", e) + logger.warning("start_write_cache error, skip saving: %s", e) with self._canceled_save_request_ids_lock: self._canceled_save_request_ids.append(req_id) return - # call manager start write - logger.debug("start_write_cache resp: %s", response) + locations = response["locations"] write_session_id = response["write_session_id"] - # check if success - if len(locations) == 0: + if not locations: try: self._manager_client.finish_write_cache({ - "trace_id": "test_test", + "trace_id": "finish_%s" % write_session_id[:8], "instance_id": self._extra_config.instance_id, "write_session_id": write_session_id, - "success_blocks": { - "bool_masks": { - "offset": 0 - } - } + "success_blocks": {"bool_masks": {"offset": 0}}, }) except Exception as e: - logger.warning("finish_write_cache failed, write_session_id: %s, error: %s", write_session_id, e) + logger.warning("finish_write_cache failed, session: %s, error: %s", + write_session_id, e) with self._canceled_save_request_ids_lock: self._canceled_save_request_ids.append(req_id) return need_block_idx = self.parse_block_mask_to_save_indices(response, target_save_num) - logger.debug("target_save_num: %s, need_block_idx: %s", target_save_num, need_block_idx) - message = CoordinateMessage(time.time(), SendBlockStartEvent(request_id=req_id, - write_session_id=write_session_id, - locations=locations)) + message = CoordinateMessage(time.time(), SendBlockStartEvent( + request_id=req_id, write_session_id=write_session_id, locations=locations)) self._coordinator_client.send(CoordinateMsgSerializer.dumps(message)) with self._waiting_to_save_requests_lock: self._waiting_to_save_requests.append(SaveRequest( - req_id, - locations, - need_block_idx, - write_session_id - )) + req_id, locations, need_block_idx, write_session_id)) def handle_canceled_save_req(self): - canceled_save_req_ids = [] with self._canceled_save_request_ids_lock: - canceled_save_req_ids = self._canceled_save_request_ids + canceled = self._canceled_save_request_ids self._canceled_save_request_ids = [] - for canceled_req_id in canceled_save_req_ids: - req = self._alive_requests[canceled_req_id] + for req_id in canceled: + req = self._alive_requests[req_id] req.sent_saving_count += 1 if (req.need_report_after_saving_finished and req.scheduled_saving_count == req.sent_saving_count): @@ -833,47 +884,37 @@ def handle_canceled_save_req(self): self._alive_requests.pop(req.req_id) def get_finished_count(self): - # only rank0 will return finished + # Only rank0 reports finished requests. return 1 def update_connector_output(self, connector_output: KVConnectorOutput): - """ - Update KVConnector state from worker-side connectors output. - - Args: - connector_output (KVConnectorOutput): the worker-side - connectors output. - """ - return - def parse_block_mask_to_save_indices(self, response: dict, target_save_num: int) -> list[int]: - # 从response中提取block_mask + def parse_block_mask_to_save_indices(self, response: dict, target_save_num: int) -> List[int]: block_mask = response.get("block_mask", {}) - save_indices = [] if "offset" in block_mask: - offset = block_mask["offset"] - for idx in range(offset, target_save_num): - save_indices.append(idx) - else: - bool_masks = block_mask.get("bool_masks", {}).get("values", []) - # 找出所有为False的索引(需要保存的block) - for idx, is_saved in enumerate(bool_masks): - if not is_saved: # False表示需要保存 - save_indices.append(idx) - - return save_indices - - def request_finished( - self, - request: "Request", - block_ids: list[int], - ) -> Tuple[bool, Optional[dict[str, Any]]]: - if request.request_id not in self._alive_requests: - logger.info("request_finished not alive request: %s", request.request_id) + return list(range(block_mask["offset"], target_save_num)) + values = block_mask.get("bool_masks", {}).get("values", []) + return [idx for idx, saved in enumerate(values) if not saved] + + # ------------------------------------------------------------------ # + # Request finish + # ------------------------------------------------------------------ # + def request_finished_all_groups( + self, request: "Request", + block_ids: Tuple[List[int], ...]) -> Tuple[bool, Optional[dict]]: + return self._finish_request(request) + + def request_finished(self, request: "Request", + block_ids: List[int]) -> Tuple[bool, Optional[dict]]: + return self._finish_request(request) + + def _finish_request(self, request: "Request") -> Tuple[bool, Optional[dict]]: + req = self._alive_requests.get(request.request_id) + if req is None: + logger.info("request_finished for unknown request: %s", request.request_id) return False, {} - req = self._alive_requests[request.request_id] extra_info = {"local_matched_token_num": req.local_matched_token_num, "remote_matched_token_num": req.remote_matched_token_num} @@ -882,8 +923,6 @@ def request_finished( self._alive_requests.pop(req.req_id) return True, extra_info - # This request still has some save requests waiting to be issued or canceled, - # delay finishing this request + # Saves still in flight; delay freeing the blocks until they land. req.need_report_after_saving_finished = True - return True, extra_info From 193a9597a69da750bb537f658e24934b8ae48614 Mon Sep 17 00:00:00 2001 From: xiaozeyu Date: Mon, 27 Jul 2026 21:36:59 +0800 Subject: [PATCH 2/7] [integration_test] add vLLM e2e KV verification for full + hybrid attention Two-phase save/load verification driven through a real KVCM manager and a real vLLM OpenAI server: phase 1 saves KV to KVCM and captures references straight from vLLM's paged cache; phase 2 loads via the connector and captures again. Captures use only vLLM's own block-table mapping, breaking the save/load symmetry so per-group translation bugs cannot cancel out. Bit-exact compare with cosine > 99.99% fallback. Scenarios: test_basic (TP=1), test_concurrent (4 reqs), test_tp (TP=2 with cross-block manager/vllm block size mapping for full-attention models). The same targets run against full-attention (Qwen2.5) and hybrid (Qwen3.5) models via KVCM_E2E_MODEL; hybrid runs enable prefix caching (mamba align mode) and restart vLLM between phases so loads come from KVCM, not the local cache. Adds a matrixed GitHub workflow (full-attention + hybrid) for self-hosted GPU runners. --- .github/workflows/test-vllm-e2e.yml | 111 ++++ integration_test/vllm_e2e/BUILD | 61 ++ integration_test/vllm_e2e/README.md | 102 +++ integration_test/vllm_e2e/e2e_lib.py | 636 +++++++++++++++++++ integration_test/vllm_e2e/test_basic.py | 27 + integration_test/vllm_e2e/test_concurrent.py | 28 + integration_test/vllm_e2e/test_connector.py | 287 +++++++++ integration_test/vllm_e2e/test_tp.py | 33 + 8 files changed, 1285 insertions(+) create mode 100644 .github/workflows/test-vllm-e2e.yml create mode 100644 integration_test/vllm_e2e/BUILD create mode 100644 integration_test/vllm_e2e/README.md create mode 100644 integration_test/vllm_e2e/e2e_lib.py create mode 100644 integration_test/vllm_e2e/test_basic.py create mode 100644 integration_test/vllm_e2e/test_concurrent.py create mode 100644 integration_test/vllm_e2e/test_connector.py create mode 100644 integration_test/vllm_e2e/test_tp.py diff --git a/.github/workflows/test-vllm-e2e.yml b/.github/workflows/test-vllm-e2e.yml new file mode 100644 index 000000000..6afc42c2e --- /dev/null +++ b/.github/workflows/test-vllm-e2e.yml @@ -0,0 +1,111 @@ +name: test-vllm-e2e +permissions: + contents: read +on: + pull_request: + branches: ["main"] + workflow_dispatch: + inputs: + runs-on: + description: "GPU runner label (needs >= 2 GPUs, e.g. A10)" + type: string + default: "gpu-a10-x2" + +concurrency: + group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} + cancel-in-progress: ${{ github.event_name == 'pull_request' }} + +jobs: + vllm-e2e: + name: vllm-e2e (${{ matrix.kind }}) + # GPU runners are self-hosted; the default label can be overridden per-repo + # via the VLLM_E2E_RUNS_ON variable or the workflow_dispatch input. + runs-on: ${{ inputs.runs-on || vars.VLLM_E2E_RUNS_ON || 'gpu-a10-x2' }} + timeout-minutes: 180 + strategy: + fail-fast: false + matrix: + include: + # Full-attention model: single FullAttentionSpec kv cache group. + - kind: full-attention + model_var: VLLM_E2E_MODEL_FULL_ATTN + # Hybrid model: MambaSpec groups + FullAttentionSpec group + # (mamba_cache_mode="align"). + - kind: hybrid-attention + model_var: VLLM_E2E_MODEL_HYBRID + steps: + - uses: actions/checkout@v4 + + - name: check_gpus + run: | + nvidia-smi --query-gpu=index,name,memory.total --format=csv,noheader + GPU_COUNT=$(nvidia-smi --query-gpu=name --format=csv,noheader | wc -l) + if [ "$GPU_COUNT" -lt 2 ]; then + echo "::error::This test needs at least 2 GPUs, found $GPU_COUNT" + exit 1 + fi + + - name: resolve_runner_config + # Self-hosted runner provisioning: a vLLM (>= 0.26.0) venv and the model + # checkouts are pre-cached on the runner and exposed via repo variables: + # VLLM_E2E_PYTHON venv python with vllm installed + # VLLM_E2E_MODEL_FULL_ATTN e.g. a local Qwen2.5-7B-Instruct checkout + # VLLM_E2E_MODEL_HYBRID e.g. a local Qwen3.5-4B checkout + env: + VENV_PY: ${{ vars.VLLM_E2E_PYTHON }} + MODEL: ${{ vars[matrix.model_var] }} + run: | + if [ ! -x "$VENV_PY" ]; then + echo "::error::VLLM_E2E_PYTHON ($VENV_PY) not found; provision the runner" + exit 1 + fi + "$VENV_PY" -c 'import vllm; v = vllm.__version__; print("vllm", v)' + if [ ! -f "$MODEL/config.json" ]; then + echo "::error::model not found at $MODEL; provision the runner" + exit 1 + fi + echo "KVCM_E2E_PYTHON=$VENV_PY" >> "$GITHUB_ENV" + echo "KVCM_E2E_MODEL=$MODEL" >> "$GITHUB_ENV" + + - name: build_binaries + run: | + set -x + bazelisk build //kv_cache_manager:kv_cache_manager_bin \ + //kv_cache_manager/client/pybind:kvcm_py_client_lib_wheel \ + //kv_cache_manager/py_connector/vllm:kvcm_vllm_connector_wheel \ + --per_file_copt='external/jsoncpp_git/.*@-Wno-error' + + - name: install_wheels + # Bazel wheel filenames contain unstamped {STABLE_*} template variables; + # read the real version from the wheel METADATA and rename before install. + run: | + set -x + mkdir -p /tmp/kvcm_whl && rm -f /tmp/kvcm_whl/*.whl + for whl in bazel-bin/kv_cache_manager/client/pybind/kvcm_py_client-*.whl \ + bazel-bin/kv_cache_manager/py_connector/vllm/kvcm_vllm_connector-*.whl; do + pkg=$(basename "$whl" | sed 's/-{STABLE.*//') + ver=$(unzip -p "$whl" "*.dist-info/METADATA" | awk '/^Version:/{print $2; exit}') + cp "$whl" "/tmp/kvcm_whl/${pkg}-${ver}-cp312-cp312-manylinux_2_32_x86_64.whl" + done + "$KVCM_E2E_PYTHON" -m pip install --no-deps --force-reinstall /tmp/kvcm_whl/*.whl || \ + uv pip install --python "$KVCM_E2E_PYTHON" --no-deps --force-reinstall /tmp/kvcm_whl/*.whl + + - name: run_e2e_tests + run: | + set -x + bazelisk test //integration_test/vllm_e2e/... \ + --cache_test_results=no --test_output=errors \ + --test_env=KVCM_E2E_PYTHON="$KVCM_E2E_PYTHON" \ + --test_env=KVCM_E2E_MODEL="$KVCM_E2E_MODEL" \ + --per_file_copt='external/jsoncpp_git/.*@-Wno-error' + + - name: upload_logs + if: failure() + uses: actions/upload-artifact@v6 + with: + name: vllm-e2e-logs-${{ matrix.kind }} + path: | + /tmp/kvcm_vllm_e2e/**/*.stdout + /tmp/kvcm_vllm_e2e/**/*.stderr + bazel-out/*-opt/testlogs/integration_test/vllm_e2e/** + if-no-files-found: ignore diff --git a/integration_test/vllm_e2e/BUILD b/integration_test/vllm_e2e/BUILD new file mode 100644 index 000000000..c025478cd --- /dev/null +++ b/integration_test/vllm_e2e/BUILD @@ -0,0 +1,61 @@ +package(default_visibility = ["//integration_test:__subpackages__"]) + +# Shared library: orchestration (manager + vLLM + driver + comparison) and the +# verifying connector injected into vLLM via kv_connector_module_path. +py_library( + name = "e2e_lib", + srcs = [ + "e2e_lib.py", + "test_connector.py", + ], + imports = ["."], + tags = ["no-remote-exec"], +) + +py_test( + name = "test_basic", + srcs = ["test_basic.py"], + data = [ + "//kv_cache_manager:kv_cache_manager_bin", + ], + imports = ["."], + tags = [ + "no-remote-exec", + "gpu", # requires 1+ GPU + "exclusive", # GPU tests must run serially to avoid CUDA OOM contention + ], + timeout = "eternal", + deps = [":e2e_lib"], +) + +py_test( + name = "test_concurrent", + srcs = ["test_concurrent.py"], + data = [ + "//kv_cache_manager:kv_cache_manager_bin", + ], + imports = ["."], + tags = [ + "no-remote-exec", + "gpu", # requires 1+ GPU + "exclusive", # GPU tests must run serially to avoid CUDA OOM contention + ], + timeout = "eternal", + deps = [":e2e_lib"], +) + +py_test( + name = "test_tp", + srcs = ["test_tp.py"], + data = [ + "//kv_cache_manager:kv_cache_manager_bin", + ], + imports = ["."], + tags = [ + "no-remote-exec", + "gpu", # requires 1+ GPU + "exclusive", # GPU tests must run serially to avoid CUDA OOM contention + ], + timeout = "eternal", + deps = [":e2e_lib"], +) diff --git a/integration_test/vllm_e2e/README.md b/integration_test/vllm_e2e/README.md new file mode 100644 index 000000000..40c82f624 --- /dev/null +++ b/integration_test/vllm_e2e/README.md @@ -0,0 +1,102 @@ +# vLLM <-> KVCM End-to-End KV Cache Verification + +End-to-end integration tests for the KVCM vLLM connector +(`kv_cache_manager/py_connector/vllm`). Each test starts a real KVCM manager +(local-file storage backend) and a real vLLM OpenAI server, drives prompts +through the OpenAI API and verifies that the KV cache data saved to / loaded +from KVCM is correct. + +Requires 1-2 GPUs and vLLM >= 0.26.0. + +## What is verified + +The connector translates between three block spaces per `kv_cache_group`: + +``` +KVCM manager block idx -> global token idx -> group logical block + (step 1, connector-only) (step 2/3, shared with vLLM) +``` + +A bug in step 1 is *symmetric*: save gathers from the wrong slots and load +scatters back to the same wrong slots, so a transport round trip alone cannot +detect it. The test breaks the symmetry with `VerifyingConnector` +(`test_connector.py`), a subclass of the production connector that +independently captures KV data from vLLM's paged cache using only vLLM's own +block-table mapping: + +1. **Phase 1** — fresh prompts: prefill -> connector saves to KVCM. The saved + token ranges are captured from the paged cache (**reference** captures). +2. **Phase 2** — same prompts + suffix: connector reports an external match and + loads from KVCM. The loaded blocks are captured (**loaded** captures). +3. The driver (`e2e_lib.py`) matches loaded captures against references by + token content and compares per layer: bit-exact preferred, cosine + similarity > 99.99% as fallback. + +## Model coverage + +The same test targets run against either model kind, selected by +`KVCM_E2E_MODEL`: + +| Kind | Example | Groups | Orchestration | +|---|---|---|---| +| Full attention | Qwen2.5-7B-Instruct | 1 `FullAttentionSpec` | prefix caching off, one server for both phases | +| Hybrid | Qwen3.5-4B | 3 `MambaSpec` + 1 `FullAttentionSpec` | prefix caching on (`mamba_cache_mode="align"`), server restarted between phases so phase 2 loads from KVCM instead of the local prefix cache | + +Hybrid specifics verified: + +* Per-group location specs (`tp{rank}_g{group}`) and per-group block tables. +* Attention groups: token-granular gather/scatter through the Triton kernel. +* Mamba/linear groups: per-block opaque state copy, where a manager block's + *last* token selects the state block (`_state_block_ids`). + +## Scenarios + +| Test | TP | Prompts | Notes | +|---|---|---|---| +| `test_basic` | 1 | 1 | Minimal save -> load round trip | +| `test_concurrent` | 1 | 4 | Concurrent requests: ReqState tracking, per-request block attribution | +| `test_tp` | 2 | 2 | TP coordination; for full-attention models also `preferred_block_size=32` != vLLM block size (16), forcing real cross-block translation | + +## Running + +Build prerequisites (from the repo root): + +```bash +bazelisk build //kv_cache_manager:kv_cache_manager_bin \ + //kv_cache_manager/client/pybind:kvcm_py_client_lib_wheel \ + //kv_cache_manager/py_connector/vllm:kvcm_vllm_connector_wheel \ + --per_file_copt='external/jsoncpp_git/.*@-Wno-error' +``` + +Install both wheels into the vLLM venv (rename them first: the Bazel output +name contains unstamped `{STABLE_*}` template variables; read the real version +from the wheel's `METADATA`). + +Run (tagged `exclusive`, so they execute serially): + +```bash +bazelisk test //integration_test/vllm_e2e/... \ + --cache_test_results=no --test_output=errors \ + --test_env=KVCM_E2E_PYTHON=/path/to/vllm-venv/bin/python \ + --test_env=KVCM_E2E_MODEL=/path/to/model \ + --per_file_copt='external/jsoncpp_git/.*@-Wno-error' +``` + +Environment variables: + +| Variable | Meaning | +|---|---| +| `KVCM_E2E_PYTHON` | Python interpreter with vLLM + both KVCM wheels installed | +| `KVCM_E2E_MODEL` | Model path; hybrid models are auto-detected from `config.json` | + +## Debugging + +Bazel's `test.log` only shows the driver's view (e.g. HTTP 500). The real +tracebacks live in the scenario workdir under `$TEST_TMPDIR`: + +``` +/kvcm_vllm_e2e// + manager/manager.stdout|stderr # KVCM manager + vllm/vllm*.stdout|stderr # vLLM (EngineCore tracebacks are here) + captures/{ref|loaded}_tp{rank}_{token_hash}.pt +``` diff --git a/integration_test/vllm_e2e/e2e_lib.py b/integration_test/vllm_e2e/e2e_lib.py new file mode 100644 index 000000000..ae7698862 --- /dev/null +++ b/integration_test/vllm_e2e/e2e_lib.py @@ -0,0 +1,636 @@ +"""Orchestration for the KVCM <-> vLLM end-to-end KV cache verification test. + +This module is imported by the Bazel ``py_test`` targets. It: + +1. Starts a KVCM manager (``kv_cache_manager_bin``) with a local-file storage + backend. +2. Starts a vLLM OpenAI server configured with the ``VerifyingConnector`` + (injected via ``kv_connector_module_path`` -- no vLLM files are modified). +3. Drives prompts through the OpenAI API in two phases and compares the + independently-captured KV data (reference from the save path vs loaded from + the load path). + +The comparison is done on KV data captured from vLLM's paged cache using vLLM's +own block-table mapping (independent of the connector's per-group translation), +which is what makes the test able to detect symmetric save/load translation +bugs. See ``test_connector.py`` for the capture-side details. + +Full-attention vs hybrid models +------------------------------- +The same test targets run against either model, selected by ``$KVCM_E2E_MODEL``: + +* Full-attention (e.g. Qwen2.5): a single ``FullAttentionSpec`` group. Prefix + caching is disabled so every phase-2 request is served through the connector + (no local prefix hit); the two phases share one server. +* Hybrid (e.g. Qwen3.5): several ``MambaSpec`` groups plus a ``FullAttentionSpec`` + group. Prefix caching must be enabled for vLLM to produce per-group block + tables (``mamba_cache_mode="align"``). Because that also populates the local + prefix cache, the vLLM server is restarted between the two phases so phase 2 + genuinely loads from KVCM instead of hitting the local cache. +""" + +import glob +import json +import logging +import os +import shutil +import socket +import subprocess +import time +import uuid +from typing import Optional + +import requests + +logger = logging.getLogger("vllm_e2e") + +MODEL_PATH = os.environ.get("KVCM_E2E_MODEL", "/root/ws/resources/models/Qwen2.5-7B-Instruct") +COSINE_THRESHOLD = 0.9999 + + +def is_hybrid_model(model_path: str) -> bool: + """Detect a hybrid (mamba/linear + full attention) model from its config.""" + try: + with open(os.path.join(model_path, "config.json")) as f: + cfg = json.load(f) + except Exception: + return False + text_cfg = cfg.get("text_config", cfg) + # Hybrid models interleave linear/mamba layers with full attention and + # expose a full_attention_interval / linear_* knob. + return ( + "full_attention_interval" in text_cfg + or "linear_conv_kernel_dim" in text_cfg + or cfg.get("model_type", "").startswith("qwen3_5") + ) + + +# --------------------------------------------------------------------------- # +# Paths / binaries +# --------------------------------------------------------------------------- # +def _runfiles_root() -> Optional[str]: + return os.environ.get("RUNFILES_DIR") or os.environ.get("TEST_SRCDIR") + + +def find_repo_root() -> str: + """Locate the KVCM repository root (works under Bazel runfiles and plain).""" + here = os.path.dirname(os.path.abspath(__file__)) + # integration_test/vllm_e2e/e2e_lib.py -> repo root is two levels up. + candidate = os.path.abspath(os.path.join(here, "..", "..")) + if os.path.exists(os.path.join(candidate, "WORKSPACE")): + return candidate + runfiles = _runfiles_root() + if runfiles: + cand = os.path.join(runfiles, "kv_cache_manager") + if os.path.exists(os.path.join(cand, "WORKSPACE")): + return cand + return candidate + + +def find_manager_binary(repo_root: str) -> str: + candidates = [ + os.path.join(repo_root, "bazel-bin/kv_cache_manager/kv_cache_manager_bin"), + os.path.join(repo_root, "bazel-out/k8-opt/bin/kv_cache_manager/kv_cache_manager_bin"), + ] + runfiles = _runfiles_root() + if runfiles: + candidates.append( + os.path.join(runfiles, "kv_cache_manager", "kv_cache_manager", + "kv_cache_manager_bin") + ) + for c in candidates: + if os.path.exists(c): + return c + raise RuntimeError( + "kv_cache_manager_bin not found; build it with: " + "bazelisk build //kv_cache_manager:kv_cache_manager_bin" + ) + + +def find_python() -> str: + return os.environ.get("KVCM_E2E_PYTHON", "/root/ws/env/global_vllm/.venv/bin/python") + + +def free_port() -> int: + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.bind(("", 0)) + return s.getsockname()[1] + + +def wait_http(url: str, timeout: float, post_body: Optional[dict] = None) -> bool: + deadline = time.time() + timeout + while time.time() < deadline: + try: + if post_body is not None: + r = requests.post(url, json=post_body, timeout=3) + else: + r = requests.get(url, timeout=3) + if r.status_code < 500: + return True + except Exception: + pass + time.sleep(1.0) + return False + + +# --------------------------------------------------------------------------- # +# KVCM manager +# --------------------------------------------------------------------------- # +class ManagerProcess: + def __init__(self, workdir: str, storage_root: str): + self.workdir = workdir + os.makedirs(workdir, exist_ok=True) + self.rpc_port = free_port() + self.http_port = free_port() + self.admin_rpc_port = free_port() + self.admin_http_port = free_port() + self.storage_root = storage_root + self.proc: Optional[subprocess.Popen] = None + self.config_path = os.path.join(workdir, "startup_config.json") + + def manager_uri(self) -> str: + return f"http://127.0.0.1:{self.http_port}" + + def _write_config(self): + cfg = { + "storage_config": { + "type": "file", + "global_unique_name": "nfs_01", + "storage_spec": { + "root_path": self.storage_root, + "key_count_per_file": 8, + }, + }, + "instance_group": { + "name": "default", + "storage_candidates": ["nfs_01"], + "global_quota_group_name": "default_quota_group", + "max_instance_count": 100, + "quota": { + "capacity": 30000000000, + "quota_config": [ + {"storage_type": "file", "capacity": 10000000000}, + {"storage_type": "hf3fs", "capacity": 10000000000}, + {"storage_type": "pace", "capacity": 10000000000}, + ], + }, + "cache_config": { + "reclaim_strategy": { + "reclaim_policy": 1, + "trigger_strategy": {"used_percentage": 0.8}, + "delay_before_delete_ms": 1000, + }, + "cache_prefer_strategy": 2, + "meta_indexer_config": { + "max_key_count": 1000000, + "mutex_shard_num": 16, + "batch_key_size": 16, + "meta_storage_backend_config": { + "storage_type": "local", + "storage_uri": "", + }, + "meta_cache_policy_config": { + "type": "LRU", + "capacity": 10000, + "cache_shard_bits": 0, + "high_pri_pool_ratio": 0.0, + }, + }, + }, + "user_data": '{"description": "vllm e2e test instance group"}', + "version": 1, + }, + } + with open(self.config_path, "w") as f: + json.dump(cfg, f, indent=2) + + def start(self, repo_root: str): + self._write_config() + binary = find_manager_binary(repo_root) + cmd = [ + binary, + "--env", f"kvcm.service.rpc_port={self.rpc_port}", + "--env", f"kvcm.service.http_port={self.http_port}", + "--env", f"kvcm.service.admin_rpc_port={self.admin_rpc_port}", + "--env", f"kvcm.service.admin_http_port={self.admin_http_port}", + "--env", f"kvcm.startup_config={self.config_path}", + "--env", "kvcm.logger.log_level=5", + ] + logger.info("starting manager: %s (cwd=%s)", " ".join(cmd), self.workdir) + self.proc = subprocess.Popen( + cmd, + cwd=self.workdir, + stdout=open(os.path.join(self.workdir, "manager.stdout"), "w"), + stderr=open(os.path.join(self.workdir, "manager.stderr"), "w"), + ) + if not wait_http( + f"{self.manager_uri()}/api/getClusterInfo", + timeout=60, + post_body={"trace_id": "probe", "instance_id": "probe"}, + ): + raise RuntimeError("manager did not become ready; see manager.stderr") + logger.info("manager ready at %s", self.manager_uri()) + + def stop(self): + if self.proc and self.proc.poll() is None: + self.proc.terminate() + try: + self.proc.wait(timeout=10) + except subprocess.TimeoutExpired: + self.proc.kill() + + +# --------------------------------------------------------------------------- # +# vLLM server +# --------------------------------------------------------------------------- # +class VllmServer: + def __init__(self, workdir: str, capture_dir: str, manager_uri: str, + tp_size: int, coordinator_base_port: int, + instance_id: str, preferred_block_size: int, + enable_prefix_caching: bool): + self.workdir = workdir + os.makedirs(workdir, exist_ok=True) + self.capture_dir = capture_dir + os.makedirs(capture_dir, exist_ok=True) + self.port = free_port() + self.manager_uri = manager_uri + self.tp_size = tp_size + self.coordinator_base_port = coordinator_base_port + self.instance_id = instance_id + self.preferred_block_size = preferred_block_size + self.enable_prefix_caching = enable_prefix_caching + self.proc: Optional[subprocess.Popen] = None + + def base_url(self) -> str: + return f"http://127.0.0.1:{self.port}" + + def start(self, repo_root: str, log_suffix: str = ""): + extra_config = { + "manager_uri": self.manager_uri, + "coordinator_base_port": self.coordinator_base_port, + "instance_group": "default", + "instance_id": self.instance_id, + "preferred_block_size": self.preferred_block_size, + "log_level": "INFO", + } + kv_transfer_config = { + "kv_connector": "VerifyingConnector", + "kv_role": "kv_both", + "kv_connector_module_path": "test_connector", + "kv_connector_extra_config": extra_config, + } + cmd = [ + find_python(), "-m", "vllm.entrypoints.openai.api_server", + "--model", MODEL_PATH, + "--served-model-name", "qwen", + "--port", str(self.port), + "--tensor-parallel-size", str(self.tp_size), + "--max-model-len", "4096", + "--gpu-memory-utilization", "0.85", + "--enforce-eager", + "--max-num-seqs", "16", + "--kv-transfer-config", json.dumps(kv_transfer_config), + ] + if self.enable_prefix_caching: + # Hybrid models need prefix caching to expose per-group block tables + # (mamba_cache_mode="align"); align mode requires chunked prefill. + cmd += ["--enable-prefix-caching", "--enable-chunked-prefill"] + else: + cmd += ["--no-enable-prefix-caching"] + env = os.environ.copy() + env["PYTHONPATH"] = os.path.dirname(os.path.abspath(__file__)) + os.pathsep + env.get("PYTHONPATH", "") + env["KVCM_E2E_CAPTURE_DIR"] = self.capture_dir + # Keep the connector's KV cache layout matching its expected + # [2, num_blocks, block_size, num_kv_heads, head_size] shape. + env.setdefault("VLLM_KV_CACHE_LAYOUT", "NHD") + # Force FlashAttention for the full-attention layers: it produces the + # [2, num_blocks, block_size, num_kv_heads, head_size] layout the + # connector expects, and avoids the flashinfer backend entirely. + env.setdefault("VLLM_ATTENTION_BACKEND", "FLASH_ATTN") + # Use the PyTorch-native sampler; the flashinfer sampler JIT-compiles + # with ninja, which is not available in the test environment. + env.setdefault("VLLM_USE_FLASHINFER_SAMPLER", "0") + # The test venv ships mismatched flashinfer / flashinfer-cubin wheels; + # skip the version check so importing vLLM's attention registry does not + # crash before the FlashAttention backend is selected. + env.setdefault("FLASHINFER_DISABLE_VERSION_CHECK", "1") + logger.info("starting vllm: %s", " ".join(cmd)) + self.proc = subprocess.Popen( + cmd, + cwd=self.workdir, + env=env, + stdout=open(os.path.join(self.workdir, f"vllm{log_suffix}.stdout"), "w"), + stderr=open(os.path.join(self.workdir, f"vllm{log_suffix}.stderr"), "w"), + ) + if not wait_http(f"{self.base_url()}/health", timeout=600): + raise RuntimeError("vllm did not become ready; see vllm.stderr") + logger.info("vllm ready at %s", self.base_url()) + + def stop(self): + if self.proc and self.proc.poll() is None: + self.proc.terminate() + try: + self.proc.wait(timeout=15) + except subprocess.TimeoutExpired: + self.proc.kill() + + +# --------------------------------------------------------------------------- # +# Request driver +# --------------------------------------------------------------------------- # +def tokenize(base_url: str, prompt: str) -> list[int]: + r = requests.post(f"{base_url}/tokenize", + json={"model": "qwen", "prompt": prompt}, timeout=60) + r.raise_for_status() + return r.json()["tokens"] + + +def wait_for_prefix_cached(manager_uri: str, instance_id: str, + token_ids: list[int], timeout: float = 120.0) -> bool: + """Poll the manager until the prefix for token_ids is cached (committed). + + Mirrors what the connector's get_num_new_matched_tokens queries + (query_type=QT_PREFIX_MATCH, block_mask offset=0 for a fresh request), so it + guarantees phase 2 will actually hit the external cache. + """ + deadline = time.time() + timeout + payload = { + "trace_id": "e2e_probe", + "token_ids": token_ids, + "instance_id": instance_id, + "query_type": "QT_PREFIX_MATCH", + "block_mask": {"offset": 0}, + } + while time.time() < deadline: + try: + r = requests.post(f"{manager_uri}/api/getCacheLocation", + json=payload, timeout=10) + if r.status_code == 200: + data = r.json() + if data.get("header", {}).get("status", {}).get("code") == "OK": + locs = data.get("locations", []) + if locs: + logger.info("prefix cached: %d location(s)", len(locs)) + return True + except Exception: + pass + time.sleep(1.0) + logger.warning("timed out waiting for prefix to be cached") + return False + + +def send_completions(base_url: str, prompts: list[str], max_tokens: int = 4, + temperature: float = 0.0) -> list[dict]: + """Send prompts concurrently and return the OpenAI responses.""" + from concurrent.futures import ThreadPoolExecutor + + client_url = f"{base_url}/v1/completions" + + def _one(prompt: str) -> dict: + payload = { + "model": "qwen", + "prompt": prompt, + "max_tokens": max_tokens, + "temperature": temperature, + } + r = requests.post(client_url, json=payload, timeout=300) + r.raise_for_status() + return r.json() + + with ThreadPoolExecutor(max_workers=max(1, len(prompts))) as ex: + return list(ex.map(_one, prompts)) + + +# --------------------------------------------------------------------------- # +# Capture comparison +# --------------------------------------------------------------------------- # +def count_captures(capture_dir: str, kind: str) -> int: + return len(glob.glob(os.path.join(capture_dir, f"{kind}_*.pt"))) + + +def wait_for_captures(capture_dir: str, kind: str, expected: int, + timeout: float = 120.0): + deadline = time.time() + timeout + while time.time() < deadline: + n = count_captures(capture_dir, kind) + if n >= expected: + logger.info("saw %d/%d %s captures", n, expected, kind) + return n + time.sleep(1.0) + n = count_captures(capture_dir, kind) + logger.warning("timed out waiting for %s captures: got %d, want %d", + kind, n, expected) + return n + + +def _cosine(a, b) -> float: + import torch + a = a.reshape(-1).float() + b = b.reshape(-1).float() + denom = (a.norm() * b.norm()).clamp_min(1e-12) + return float((a @ b) / denom) + + +def compare_captures(capture_dir: str, tp_size: int) -> dict: + """Compare loaded captures against reference captures. + + Every block that was *loaded* from KVCM must correspond to a *reference* + capture (same tp rank + token content) with matching KV data. The direction + matters: saves are incremental, so some saved blocks may legitimately not be + reloaded (e.g. the tokenization boundary block) -- but every loaded block + must match something that was saved. + + Returns a report dict; the caller asserts on it. + """ + import torch + + refs = {} + loaded = {} + for path in glob.glob(os.path.join(capture_dir, "*.pt")): + name = os.path.basename(path)[:-3] # strip .pt + parts = name.split("_") + kind, tp, token_hash = parts[0], parts[1], "_".join(parts[2:]) + key = (tp, token_hash) + (refs if kind == "ref" else loaded)[key] = path + + report = { + "num_refs": len(refs), + "num_loaded": len(loaded), + "matched": 0, + "bit_exact": 0, + "cosine_pass": 0, + "failures": [], + "loaded_without_ref": [], + } + + for key, loaded_path in sorted(loaded.items()): + if key not in refs: + report["loaded_without_ref"].append(key) + continue + ref = torch.load(refs[key], map_location="cpu", weights_only=True) + got = torch.load(loaded_path, map_location="cpu", weights_only=True) + + assert ref["token_ids"] == got["token_ids"], f"token id mismatch for {key}" + + all_bit_exact = True + worst_cosine = 1.0 + for layer_name, ref_kv in ref["kv"].items(): + got_kv = got["kv"][layer_name] + # Attention groups are a single Tensor; mamba/linear/gdn groups are a + # list[Tensor] (e.g. [conv_state, ssm_state]). Compare uniformly. + if isinstance(ref_kv, (list, tuple)): + ref_parts = list(ref_kv) + got_parts = list(got_kv) + assert len(ref_parts) == len(got_parts), ( + f"state count mismatch {layer_name}: " + f"{len(ref_parts)} vs {len(got_parts)}" + ) + else: + ref_parts = [ref_kv] + got_parts = [got_kv] + + for si, (ref_t, got_t) in enumerate(zip(ref_parts, got_parts)): + assert ref_t.shape == got_t.shape, ( + f"shape mismatch {layer_name}[{si}]: {ref_t.shape} vs {got_t.shape}" + ) + if not torch.equal(ref_t, got_t): + all_bit_exact = False + cos = _cosine(ref_t, got_t) + worst_cosine = min(worst_cosine, cos) + if cos < COSINE_THRESHOLD: + report["failures"].append({ + "key": key, + "layer": f"{layer_name}[{si}]", + "cosine": cos, + }) + + report["matched"] += 1 + if all_bit_exact: + report["bit_exact"] += 1 + else: + report["cosine_pass"] += 1 + logger.warning("capture %s not bit-exact (worst cosine=%.6f)", + key, worst_cosine) + + return report + + +def assert_report_ok(report: dict): + problems = [] + if report["loaded_without_ref"]: + problems.append( + f"loaded captures with no matching reference: {report['loaded_without_ref']}" + ) + if report["failures"]: + problems.append(f"cosine failures: {report['failures']}") + if report["matched"] == 0: + problems.append("no loaded captures were matched against references") + if problems: + raise AssertionError("KV verification failed: " + "; ".join(problems)) + logger.info( + "KV verification OK: matched=%d bit_exact=%d cosine_pass=%d (refs=%d loaded=%d)", + report["matched"], report["bit_exact"], report["cosine_pass"], + report["num_refs"], report["num_loaded"], + ) + + +# --------------------------------------------------------------------------- # +# Scenario runner +# --------------------------------------------------------------------------- # +def run_e2e(scenario: str, tp_size: int, num_prompts: int, + preferred_block_size: int): + """Run one full save-then-load verification scenario. + + Full-attention models: prefix caching off, one server across both phases. + Hybrid models: prefix caching on (align mode -> per-group block tables), the + vLLM server is restarted between phases so phase 2 loads from KVCM instead of + hitting the local prefix cache. + """ + import torch # noqa: F401 (ensure torch importable early for clear errors) + + hybrid = is_hybrid_model(MODEL_PATH) + # Hybrid mamba state is per scheduler block, so preferred_block_size can only + # differ from the vLLM block size for pure-attention models. + if hybrid: + preferred_block_size = 0 + + repo_root = find_repo_root() + scratch_root = os.environ.get("TEST_TMPDIR") or os.environ.get("TMPDIR") or "/tmp" + base_workdir = os.path.join(scratch_root, "kvcm_vllm_e2e", scenario) + if os.path.exists(base_workdir): + shutil.rmtree(base_workdir) + storage_root = os.path.join(base_workdir, "nfs") + manager_dir = os.path.join(base_workdir, "manager") + vllm_dir = os.path.join(base_workdir, "vllm") + capture_dir = os.path.join(base_workdir, "captures") + os.makedirs(storage_root, exist_ok=True) + + instance_id = f"e2e-{scenario}-{uuid.uuid4().hex[:8]}" + logger.info("scenario=%s model=%s hybrid=%s tp=%d prompts=%d preferred_bs=%d", + scenario, MODEL_PATH, hybrid, tp_size, num_prompts, preferred_block_size) + + manager = ManagerProcess(manager_dir, storage_root) + + def make_server(coordinator_port): + return VllmServer( + vllm_dir, capture_dir, manager.manager_uri(), tp_size, + coordinator_base_port=coordinator_port, + instance_id=instance_id, + preferred_block_size=preferred_block_size, + enable_prefix_caching=hybrid, + ) + + vllm = make_server(free_port()) + + try: + manager.start(repo_root) + vllm.start(repo_root, log_suffix="" if not hybrid else "_p1") + + # Distinct, deterministic prompts. Each sentence carries a unique counter + # so every manager block has unique token content -- this avoids hash + # collisions between blocks with identical text but different KV (RoPE is + # position-dependent). Long enough to span several manager blocks so the + # translation layer is actually exercised. + base_prompts = [ + f"Prompt number {i}. " + " ".join( + f"Sentence {j} of prompt {i} has value {j * 7 + i * 131}." + for j in range(40) + ) + for i in range(num_prompts) + ] + suffixes = [f" Now answer question {i}: what is 2+2?" for i in range(num_prompts)] + + # ---- Phase 1: fresh prefill -> connector saves -> reference capture. + logger.info("phase 1: sending %d fresh prompts", num_prompts) + send_completions(vllm.base_url(), base_prompts) + wait_for_captures(capture_dir, "ref", expected=num_prompts, timeout=180) + + # The save is committed to the manager asynchronously after the ref + # capture (which fires when the save is submitted). Wait until the manager + # actually has the prefix, otherwise phase 2 would find no match. + phase2_prompts = [p + s for p, s in zip(base_prompts, suffixes)] + for p in base_prompts: + toks = tokenize(vllm.base_url(), p) + if not wait_for_prefix_cached(manager.manager_uri(), instance_id, toks): + raise AssertionError("save was not committed to the manager in time") + + # Hybrid models keep prefix caching on, which also populates the local + # prefix cache; restart vLLM so phase 2 loads from KVCM, not locally. + if hybrid: + logger.info("restarting vLLM before phase 2 (clear local prefix cache)") + vllm.stop() + vllm = make_server(free_port()) + vllm.start(repo_root, log_suffix="_p2") + + # ---- Phase 2: same prefix + suffix -> connector loads -> loaded capture. + logger.info("phase 2: sending %d prefix+suffix prompts", num_prompts) + send_completions(vllm.base_url(), phase2_prompts) + wait_for_captures(capture_dir, "loaded", expected=num_prompts, timeout=180) + + report = compare_captures(capture_dir, tp_size) + assert_report_ok(report) + logger.info("scenario %s PASSED: %s", scenario, json.dumps( + {k: v for k, v in report.items() if k != "failures"}, default=str)) + finally: + vllm.stop() + manager.stop() diff --git a/integration_test/vllm_e2e/test_basic.py b/integration_test/vllm_e2e/test_basic.py new file mode 100644 index 000000000..a409d12f6 --- /dev/null +++ b/integration_test/vllm_e2e/test_basic.py @@ -0,0 +1,27 @@ +"""test_basic: single-request save/load KV verification (TP=1). + +Sends one prompt (prefill + save -> reference capture), then the same prompt +with a suffix (load + prefill -> loaded capture), and verifies the prefix KV +data matches (bit-exact preferred, cosine > 99.99% as fallback). + +Works for both full-attention and hybrid models (selected via KVCM_E2E_MODEL); +see e2e_lib.run_e2e for the per-model orchestration differences. +""" + +import unittest + +from e2e_lib import run_e2e + + +class TestBasic(unittest.TestCase): + def test_basic(self): + run_e2e( + scenario="basic", + tp_size=1, + num_prompts=1, + preferred_block_size=0, # manager block size == vllm block size + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/integration_test/vllm_e2e/test_concurrent.py b/integration_test/vllm_e2e/test_concurrent.py new file mode 100644 index 000000000..4ea7501de --- /dev/null +++ b/integration_test/vllm_e2e/test_concurrent.py @@ -0,0 +1,28 @@ +"""test_concurrent: multiple concurrent requests save/load KV verification. + +Sends several distinct prompts concurrently (all prefill + save -> reference +captures), then the same prompts each with their own suffix concurrently (all +load + prefill -> loaded captures), and verifies each request's prefix KV data +matches. This exercises ReqState tracking, per-request block attribution and +async task races. + +Works for both full-attention and hybrid models (selected via KVCM_E2E_MODEL). +""" + +import unittest + +from e2e_lib import run_e2e + + +class TestConcurrent(unittest.TestCase): + def test_concurrent(self): + run_e2e( + scenario="concurrent", + tp_size=1, + num_prompts=4, + preferred_block_size=0, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/integration_test/vllm_e2e/test_connector.py b/integration_test/vllm_e2e/test_connector.py new file mode 100644 index 000000000..42d71c0df --- /dev/null +++ b/integration_test/vllm_e2e/test_connector.py @@ -0,0 +1,287 @@ +"""A verification wrapper around the production KVCM vLLM connector. + +This connector is injected via vLLM's ``kv_connector_module_path`` and subclasses +the production ``TairKvCacheConnector`` without modifying it. Its purpose is to +independently capture the KV data that lives in vLLM's paged KV cache so that the +test driver can verify the connector's save/load translation layer. + +Why this catches translation bugs +--------------------------------- +The production connector is built around *per-group* transfer. Every +``kv_cache_group`` (a ``FullAttentionSpec`` group for pure-attention models, or +several ``MambaSpec`` groups plus one ``FullAttentionSpec`` group for hybrid +models) is a self-contained transfer unit with its own block table and its own +translation: + + KVCM manager block idx -> global token idx -> group logical block + (step 1, connector-only) (step 2/3, shared with vLLM) + +Step 1 is connector-only logic. A bug there makes *save* gather from the wrong +physical slots and *load* scatter to the wrong physical slots. Because save and +load share the same translation, a transport-level round trip still "matches" +(the bug is symmetric). + +To break the symmetry we capture KV data using ONLY the position -> physical +slot mapping that vLLM itself uses (its slot_mapping kernel), a pure function of +the group's block table and the token position. This reference is independent of +the connector's step-1 logic, so a step-1 bug makes the captured data diverge +from what the connector saved/loaded. + +Group kinds +----------- +* Attention groups -> ``torch.Tensor`` of shape + ``[2, num_blocks, kernel_block_size, num_kv_heads, head_size]``. Captured + per-token using the three-tier mapping (group logical block -> kernel physical + block) because the scheduler's group block size may exceed the kernel block + size. +* Mamba/linear/gdn groups -> ``list[Tensor]`` (e.g. ``[conv_state, ssm_state]``). + The state is stored **per group block**; we capture the whole state slice for + the group block that each manager block maps to (mirroring the connector's + ``_state_block_ids``: a manager block's *last* token selects the block). + +Capture points +-------------- +* Reference (save path): in ``wait_for_save`` we read the KV of the saved token + range straight out of the paged cache (the forward pass has completed and the + slots are not modified by the parent's async gather). +* Loaded (load path): loads are async, so the load step has no forward pass and + the worker does not yet know the request's token ids. We record the load's + per-group block tables in ``start_load_kv`` and emit the capture in a later + ``wait_for_save`` once the token ids have arrived. The loaded KV persists in + the paged cache (its blocks are allocated to the request). + +Captures are written to ``$KVCM_E2E_CAPTURE_DIR`` as ``.pt`` files named +``{ref|loaded}_tp{rank}_{token_hash}.pt`` so the out-of-process driver can match +reference vs loaded by content (the captured token ids). +""" + +import hashlib +import os +import threading +import typing + +import torch + +from kv_cache_manager.py_connector.common.logger import logger +from kv_cache_manager.py_connector.vllm.metadata import TairKvCacheConnectorMetadata +from kv_cache_manager.py_connector.vllm.v1_connector import TairKvCacheConnector + +CAPTURE_DIR_ENV = "KVCM_E2E_CAPTURE_DIR" + + +class VerifyingConnector(TairKvCacheConnector): + """Production connector + independent per-group KV capture for e2e.""" + + # ------------------------------------------------------------------ # + # Setup + # ------------------------------------------------------------------ # + def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): + super().register_kv_caches(kv_caches) + + self._capture_dir = os.environ.get(CAPTURE_DIR_ENV, "") + if self._capture_dir: + os.makedirs(self._capture_dir, exist_ok=True) + + # Snapshot the static per-group description into a capture-friendly form. + # Each entry: (group_idx, is_attention, layer_names, group_block_size, + # kernel_block_size). kernel_block_size is read straight off the tensor. + self._cap_groups = [] + for meta in self._group_metas: + if meta.is_attention: + ref = kv_caches[meta.layer_names[0]] + kernel_bs = ref.shape[2] + else: + kernel_bs = 0 + self._cap_groups.append( + (meta.group_idx, meta.is_attention, list(meta.layer_names), + meta.block_size, kernel_bs)) + + # Track completion of async load scatters. The parent's load task already + # CPU-synchronizes its own scatter before reporting the task result, so a + # threading.Event set from the done callback is sufficient to know the + # scatter is globally visible. + self._load_done_events: dict[str, list[threading.Event]] = {} + self._load_events_lock = threading.Lock() + + # Loads are async and their step has no forward pass, so the worker does + # not yet have the request's token ids. Record the load's per-group block + # tables here and emit the capture once the token ids arrive. + # req_id -> (manager_block_idxes, block_ids_per_group) + self._pending_loaded: dict[str, tuple[list, list]] = {} + + orig_factory = self._data_transfer.create_load_done_callback + + def tracking_factory(req_id, *args, **kwargs): + orig_cb = orig_factory(req_id, *args, **kwargs) + evt = threading.Event() + with self._load_events_lock: + self._load_done_events.setdefault(req_id, []).append(evt) + + def cb(task_results): + try: + orig_cb(task_results) + finally: + evt.set() + + return cb + + self._data_transfer.create_load_done_callback = tracking_factory + logger.warning( + "VerifyingConnector enabled, capture_dir=%s tp_rank=%s groups=%s " + "vllm_bs=%s manager_bs=%s", + self._capture_dir, self._tp_rank, + [(g[0], "attn" if g[1] else "state", g[3], g[4]) for g in self._cap_groups], + self._vllm_block_size, self._manager_block_size, + ) + + # ------------------------------------------------------------------ # + # Load hook: record pending loaded captures + # ------------------------------------------------------------------ # + def start_load_kv(self, forward_context, **kwargs) -> None: + meta = typing.cast(TairKvCacheConnectorMetadata, self._get_connector_metadata()) + load_reqs = [ + (lr.req_id, list(lr.manager_block_idxes), + [list(b) for b in lr.all_block_ids]) + for lr in meta.to_load_requests + if lr.all_block_ids and lr.need_load_locations + ] + + super().start_load_kv(forward_context, **kwargs) + + if getattr(self, "_capture_dir", "") and load_reqs: + for req_id, mbis, bpg in load_reqs: + self._pending_loaded[req_id] = (mbis, bpg) + logger.warning( + "VerifyingConnector recorded %d pending loaded capture(s)", + len(load_reqs)) + + # ------------------------------------------------------------------ # + # Save hook: reference captures + emit pending loaded captures + # ------------------------------------------------------------------ # + def wait_for_save(self): + meta = typing.cast(TairKvCacheConnectorMetadata, self._get_connector_metadata()) + + if getattr(self, "_capture_dir", "") and getattr(self, "_kv_caches", None): + try: + self._capture_refs(meta) + self._capture_pending_loaded() + except Exception as e: # never break inference for a capture error + logger.warning("VerifyingConnector capture failed: %s", e, exc_info=True) + + super().wait_for_save() + + def _capture_refs(self, meta: TairKvCacheConnectorMetadata): + if not meta.to_save_requests: + return + # Make all forward-pass KV writes visible before reading the paged cache. + self._device_mod.synchronize() + for save_req in meta.to_save_requests: + req = self._alive_requests.get(save_req.req_id) + if req is None or not req.block_ids_per_group: + continue + self._capture_range( + kind="ref", + token_ids=req.token_ids, + block_ids_per_group=req.block_ids_per_group, + manager_block_idxes=save_req.manager_block_idxes, + ) + + def _capture_pending_loaded(self): + if not self._pending_loaded: + return + done = [] + for req_id, (mbis, bpg) in self._pending_loaded.items(): + req = self._alive_requests.get(req_id) + if req is None: + # token ids have not arrived on this worker yet; wait for a + # later step in which the request is scheduled. + continue + with self._load_events_lock: + evts = list(self._load_done_events.get(req_id, [])) + for evt in evts: + evt.wait(timeout=120) + self._device_mod.synchronize() + self._capture_range( + kind="loaded", + token_ids=req.token_ids, + block_ids_per_group=bpg, + manager_block_idxes=mbis, + ) + done.append(req_id) + for req_id in done: + del self._pending_loaded[req_id] + + # ------------------------------------------------------------------ # + # Capture helpers + # ------------------------------------------------------------------ # + def _capture_range(self, kind, token_ids, block_ids_per_group, manager_block_idxes): + if not manager_block_idxes or not block_ids_per_group: + return + # One record per manager block. Saves are batched incrementally while + # loads arrive all-at-once, so per-block records let the driver match + # reference vs loaded captures by each block's token content. + for b in manager_block_idxes: + self._capture_block(kind, token_ids, block_ids_per_group, b) + + def _attn_token_slot(self, pos, block_table, group_bs, kernel_bs): + """Map a global token position to its flat slot in an attention group. + + Mirrors vLLM's own slot_mapping kernel expressed with the three-tier + block hierarchy (group logical block -> kernel physical block). Works for + pure-attention groups (group_bs == kernel_bs, ratio 1) and hybrid + attention groups (group block larger than kernel block). Independent of + the connector's step-1 (manager-block) logic, which is what we verify. + """ + ratio = group_bs // kernel_bs + logical = pos // group_bs + off = pos % group_bs + physical = block_table[logical] * ratio + off // kernel_bs + return physical * kernel_bs + off % kernel_bs + + def _capture_block(self, kind, token_ids, block_ids_per_group, manager_block_idx): + mbs = self._manager_block_size + + # Global token positions covered by this manager block. + positions = list(range(manager_block_idx * mbs, (manager_block_idx + 1) * mbs)) + if positions[-1] >= len(token_ids): + positions = [p for p in positions if p < len(token_ids)] + if not positions: + return + + captured_token_ids = [token_ids[p] for p in positions] + kv_by_layer = {} + + for group_idx, is_attention, layer_names, group_bs, kernel_bs in self._cap_groups: + block_table = block_ids_per_group[group_idx] + if is_attention: + slots = [self._attn_token_slot(p, block_table, group_bs, kernel_bs) + for p in positions] + slot_tensor = torch.tensor(slots, dtype=torch.long, device=self._device) + for layer_name in layer_names: + kv_cache = self._kv_caches[layer_name] + # vLLM >= 0.26.0: (num_blocks, num_kv_heads, kernel_bs, + # 2*head_size) packed, NHD memory order is token-major. Flatten + # the (block, token) dims and gather the whole per-token vector + # (K and V packed) -- the packing is opaque to verification. + per_token = kv_cache.shape[1] * kv_cache.shape[3] + flat = kv_cache.permute(0, 2, 1, 3).reshape(-1, per_token) + gathered = flat[slot_tensor, :].contiguous() # [n_tok, per_token] + kv_by_layer[layer_name] = gathered.cpu() + kv_by_layer[layer_name] = gathered.cpu() + else: + # State stored once per group block; the manager block's last + # token selects the block (mirrors _state_block_ids). + logical = ((manager_block_idx + 1) * mbs - 1) // group_bs + block_id = block_table[logical] + for layer_name in layer_names: + states = self._kv_caches[layer_name] # list[Tensor] + kv_by_layer[layer_name] = [s[block_id].detach().cpu() for s in states] + + token_hash = hashlib.sha256( + torch.tensor(captured_token_ids, dtype=torch.int64).numpy().tobytes() + ).hexdigest()[:16] + path = os.path.join(self._capture_dir, f"{kind}_tp{self._tp_rank}_{token_hash}.pt") + torch.save({"token_ids": captured_token_ids, "kv": kv_by_layer}, path) + logger.warning( + "VerifyingConnector captured %s block=%d tokens=%d..%d tp=%s -> %s", + kind, manager_block_idx, positions[0], positions[-1], self._tp_rank, path) diff --git a/integration_test/vllm_e2e/test_tp.py b/integration_test/vllm_e2e/test_tp.py new file mode 100644 index 000000000..a44c572f4 --- /dev/null +++ b/integration_test/vllm_e2e/test_tp.py @@ -0,0 +1,33 @@ +"""test_tp: TP=2 save/load KV verification with a non-trivial block translation. + +Runs the save/load verification under tensor parallelism (TP=2), where each rank +has an independent forward context, slot mapping and capture, and the connector's +ZMQ-based TP coordination is fully exercised. + +For full-attention models it also sets ``preferred_block_size=32`` while vLLM +uses its default block size (16), forcing the connector's manager-block <-> +group-block translation (``_attn_token_indices``) to do real cross-block +mapping -- the code path most prone to symmetric save/load bugs. For hybrid +models the manager block size is pinned to the scheduler block size (mamba state +is per scheduler block), so run_e2e ignores preferred_block_size there. + +Works for both full-attention and hybrid models (selected via KVCM_E2E_MODEL). +""" + +import unittest + +from e2e_lib import run_e2e + + +class TestTp(unittest.TestCase): + def test_tp(self): + run_e2e( + scenario="tp", + tp_size=2, + num_prompts=2, + preferred_block_size=32, # != vllm block size (16) for full-attn models + ) + + +if __name__ == "__main__": + unittest.main() From 593c4e6c22b4eea10d4cf00a710a616f5196ad61 Mon Sep 17 00:00:00 2001 From: xiaozeyu Date: Tue, 28 Jul 2026 16:27:24 +0800 Subject: [PATCH 3/7] [py_connector] cap full-prompt external hit and fix canceled-save KeyError Two fixes in the vLLM connector scheduler path: 1. get_num_new_matched_tokens returned the raw external match count without capping it below the prompt length. This connector loads synchronously (load_kv_async=False), so vLLM schedules num_tokens - num_computed_tokens new tokens and asserts that count is > 0 (vllm 0.26.0 v1/core/sched/scheduler.py waiting-queue loop). A prompt whose token count is an exact multiple of the manager block size with all blocks externally cached made the count 0 and crashed the engine. Drop trailing matched blocks until at least one token remains to recompute, mirroring the fallback in vLLM's own SharedStorage/NIXL connectors. 2. handle_canceled_save_req indexed _alive_requests[req_id] directly, but cancellations arrive from http_executor threads and can race request teardown; use .get() with a warning and skip. Covered by kv_cache_manager/py_connector/test/test_scheduler_state.py and the integration_test/vllm_e2e test_full_hit scenario. --- .../py_connector/vllm/v1_connector.py | 15 ++++++++++++++- 1 file changed, 14 insertions(+), 1 deletion(-) diff --git a/kv_cache_manager/py_connector/vllm/v1_connector.py b/kv_cache_manager/py_connector/vllm/v1_connector.py index d4acd634d..baa8b706a 100644 --- a/kv_cache_manager/py_connector/vllm/v1_connector.py +++ b/kv_cache_manager/py_connector/vllm/v1_connector.py @@ -709,6 +709,14 @@ def get_num_new_matched_tokens(self, request: "Request", return None, False new_matched_count = len(need_load_locations) * self._manager_block_size + # This connector loads synchronously (load_kv_async=False), so vLLM will + # schedule num_tokens - num_computed_tokens new tokens and asserts that + # count is > 0 (vllm/v1/core/sched/scheduler.py). If the whole prompt is + # externally cached, drop trailing blocks so at least one token is + # recomputed locally. + while new_matched_count and num_computed_tokens + new_matched_count >= request.num_tokens: + need_load_locations = need_load_locations[:-1] + new_matched_count -= self._manager_block_size total_remote_blocks = computed_blocks + len(need_load_locations) logger.info("req:%s matched %d external tokens", request.request_id, new_matched_count) @@ -876,7 +884,12 @@ def handle_canceled_save_req(self): canceled = self._canceled_save_request_ids self._canceled_save_request_ids = [] for req_id in canceled: - req = self._alive_requests[req_id] + # Cancellations come from http_executor threads; the request may + # already have been finished and removed by the scheduler loop. + req = self._alive_requests.get(req_id) + if req is None: + logger.warning("canceled save for unknown request %s, skip", req_id) + continue req.sent_saving_count += 1 if (req.need_report_after_saving_finished and req.scheduled_saving_count == req.sent_saving_count): From 889796b6ee79c4b85589764e4a2beedd49c629bd Mon Sep 17 00:00:00 2001 From: xiaozeyu Date: Tue, 28 Jul 2026 16:27:37 +0800 Subject: [PATCH 4/7] [py_connector] add unit tests for translation, transfer results and scheduler state Previously zero unit coverage on the connector's pure logic. New Bazel py_tests under py_connector/test (vllm_stubs.py registers lightweight vLLM / pybind stand-ins in sys.modules so v1_connector imports without a GPU or the compiled client): * test_block_translation: _attn_token_indices / _state_block_ids against an independent brute-force reference, parameterized over ratio=1, ratio>1 and manager_bs != group_bs, plus hand-computed examples. * test_data_transfer_results: MultiResult ordering (in-order, out-of-order, concurrent) and the save/load done callbacks' stride-AND merge, which pins the implicit group-major submission-order contract of _submit_group_tasks; includes the hybrid report_failures=False branch. * test_scheduler_state: get_num_new_matched_tokens (incl. the full-hit cap), parse_block_mask_to_save_indices (offset and bool_masks), _parse_groups (full-attn, hybrid, eagle skip, unsupported spec), and the build_connector_meta state machine (new request, cached deltas with new_block_ids None/non-None, preemption resume via both resumed_req_ids and legacy resumed_from_preemption, save-threshold growth, both request_finished paths, canceled-save races). * test/kernel/test_strided_gather_scatter (GPU): the strided kernel path (block_stride/local_block_size incl. padded pages) added in fc896911 had no coverage; checked element-wise against naive torch indexing, plus a roundtrip and a padding-untouched sentinel check. vllm/BUILD: expose vllm_connector to py_connector subpackages for the tests. --- kv_cache_manager/py_connector/test/BUILD | 36 ++ .../py_connector/test/kernel/BUILD | 18 + .../kernel/test_strided_gather_scatter.py | 186 +++++++++ .../test/test_block_translation.py | 125 ++++++ .../test/test_data_transfer_results.py | 136 ++++++ .../py_connector/test/test_scheduler_state.py | 391 ++++++++++++++++++ .../py_connector/test/vllm_stubs.py | 133 ++++++ kv_cache_manager/py_connector/vllm/BUILD | 1 + 8 files changed, 1026 insertions(+) create mode 100644 kv_cache_manager/py_connector/test/BUILD create mode 100644 kv_cache_manager/py_connector/test/kernel/BUILD create mode 100644 kv_cache_manager/py_connector/test/kernel/test_strided_gather_scatter.py create mode 100644 kv_cache_manager/py_connector/test/test_block_translation.py create mode 100644 kv_cache_manager/py_connector/test/test_data_transfer_results.py create mode 100644 kv_cache_manager/py_connector/test/test_scheduler_state.py create mode 100644 kv_cache_manager/py_connector/test/vllm_stubs.py diff --git a/kv_cache_manager/py_connector/test/BUILD b/kv_cache_manager/py_connector/test/BUILD new file mode 100644 index 000000000..4d11111b5 --- /dev/null +++ b/kv_cache_manager/py_connector/test/BUILD @@ -0,0 +1,36 @@ +load("@rules_python//python:py_library.bzl", "py_library") +load("@rules_python//python:py_test.bzl", "py_test") + +# Stubs that make v1_connector importable without vLLM / CUDA / the compiled +# kvcm_py_client. Tests import this module before anything under vllm/. +py_library( + name = "vllm_stubs", + srcs = [ + "__init__.py", + "vllm_stubs.py", + ], + deps = [ + "//kv_cache_manager/py_connector/vllm:vllm_connector", + ], +) + +py_test( + name = "test_block_translation", + srcs = ["test_block_translation.py"], + tags = ["no-remote-exec"], + deps = [":vllm_stubs"], +) + +py_test( + name = "test_data_transfer_results", + srcs = ["test_data_transfer_results.py"], + tags = ["no-remote-exec"], + deps = [":vllm_stubs"], +) + +py_test( + name = "test_scheduler_state", + srcs = ["test_scheduler_state.py"], + tags = ["no-remote-exec"], + deps = [":vllm_stubs"], +) diff --git a/kv_cache_manager/py_connector/test/kernel/BUILD b/kv_cache_manager/py_connector/test/kernel/BUILD new file mode 100644 index 000000000..bbea02aeb --- /dev/null +++ b/kv_cache_manager/py_connector/test/kernel/BUILD @@ -0,0 +1,18 @@ +load("@rules_python//python:py_test.bzl", "py_test") + +py_test( + name = "test_strided_gather_scatter", + srcs = [ + "__init__.py", + "test_strided_gather_scatter.py", + ], + main = "test_strided_gather_scatter.py", + tags = [ + "no-remote-exec", + "gpu", # requires 1 GPU + "exclusive", # GPU tests run serially to avoid CUDA contention + ], + deps = [ + "//kv_cache_manager/py_connector/kernel", + ], +) diff --git a/kv_cache_manager/py_connector/test/kernel/test_strided_gather_scatter.py b/kv_cache_manager/py_connector/test/kernel/test_strided_gather_scatter.py new file mode 100644 index 000000000..b1ca1caa2 --- /dev/null +++ b/kv_cache_manager/py_connector/test/kernel/test_strided_gather_scatter.py @@ -0,0 +1,186 @@ +"""GPU tests for the strided path of the batch gather/scatter Triton kernel. + +The flat path (block_stride=0) is covered by test_batch_gather_scatter.py. +Here we cover the strided path added for vLLM's paged layout, where the flat +token index is decomposed as (kv_block, token_in_block) and the block starts +``block_stride`` elements apart -- including padded pages where +``block_stride > local_block_size * dims_per_token`` leaves a gap between +blocks that must be skipped, not walked. + +Every case is checked element-wise against a naive torch reference that +performs the same (kv_block, token) decomposition with plain indexing. +""" + +import unittest + +import torch + +from kv_cache_manager.py_connector.kernel.batch_gather_scatter_helper import ( + batch_gather_kv_caches, + batch_scatter_kv_caches, +) + + +def _make_paged_caches(num_layers, num_blocks, local_block_size, dims_per_token, + pad_tokens, device, dtype, fill_random=True): + """Per-layer paged caches shaped (num_blocks, padded_tokens, dims) where + padded_tokens = local_block_size + pad_tokens. block_stride (in elements) + is padded_tokens * dims_per_token.""" + caches = [] + for _ in range(num_layers): + t = torch.randn(num_blocks, local_block_size + pad_tokens, dims_per_token, + device=device, dtype=dtype) if fill_random else \ + torch.zeros(num_blocks, local_block_size + pad_tokens, dims_per_token, + device=device, dtype=dtype) + caches.append(t) + return caches + + +def _ref_slot(cache, flat_token_idx, local_block_size): + blk = flat_token_idx // local_block_size + tok = flat_token_idx % local_block_size + return cache[blk, tok, :] + + +class TestStridedGatherScatter(unittest.TestCase): + # (local_block_size, pad_tokens, tokens_per_manager_block) + CASES = [ + (16, 0, 16), # strided == flat geometry (stride still exercised) + (16, 4, 16), # padded pages: gap between blocks + (64, 0, 528), # hybrid attention: manager block spans many kv blocks + (64, 8, 48), # padded + manager block not aligned to kv block + ] + + def setUp(self): + if not torch.cuda.is_available(): + self.skipTest("requires a GPU") + torch.manual_seed(7) + self.device = "cuda" + self.dtype = torch.bfloat16 + self.num_layers = 3 + self.dims = 128 + self.num_kv_blocks = 64 + + def _indices(self, num_manager_blocks, tokens_per_block, local_block_size): + total_tokens = self.num_kv_blocks * local_block_size + need = num_manager_blocks * tokens_per_block + assert need <= total_tokens, "test setup: not enough kv slots" + perm = torch.randperm(total_tokens)[:need] + return perm.tolist() + + def test_gather_strided_matches_reference(self): + for local_bs, pad, tokens_per_block in self.CASES: + with self.subTest(local_bs=local_bs, pad=pad, tpb=tokens_per_block): + caches = _make_paged_caches( + self.num_layers, self.num_kv_blocks, local_bs, self.dims, + pad, self.device, self.dtype) + block_stride = caches[0].stride(0) + self.assertEqual(block_stride, (local_bs + pad) * self.dims) + ptrs = torch.tensor([c.data_ptr() for c in caches], + device=self.device, dtype=torch.int64) + num_mb = 4 + token_indices = self._indices(num_mb, tokens_per_block, local_bs) + dst_block_indices = [2, 0, 3, 1] + dst = torch.zeros(num_mb, self.num_layers, tokens_per_block, + self.dims, device="cpu", dtype=self.dtype, + pin_memory=True) + batch_gather_kv_caches( + ptrs, dst, token_indices, dst_block_indices, + tokens_per_block, self.dims, + block_stride=block_stride, local_block_size=local_bs) + torch.cuda.synchronize() + + caches_cpu = [c.cpu() for c in caches] + for mb in range(num_mb): + for pos in range(tokens_per_block): + flat_idx = token_indices[mb * tokens_per_block + pos] + for layer in range(self.num_layers): + want = _ref_slot(caches_cpu[layer], flat_idx, local_bs) + got = dst[dst_block_indices[mb], layer, pos, :] + torch.testing.assert_close( + got, want, + msg=f"gather mismatch mb={mb} pos={pos} " + f"layer={layer} flat={flat_idx}") + + def test_scatter_strided_matches_reference(self): + for local_bs, pad, tokens_per_block in self.CASES: + with self.subTest(local_bs=local_bs, pad=pad, tpb=tokens_per_block): + caches = _make_paged_caches( + self.num_layers, self.num_kv_blocks, local_bs, self.dims, + pad, self.device, self.dtype, fill_random=False) + # Sentinel in the padding region: scatter must never touch it. + sentinel = 123.0 + if pad: + for c in caches: + c[:, local_bs:, :] = sentinel + block_stride = caches[0].stride(0) + ptrs = torch.tensor([c.data_ptr() for c in caches], + device=self.device, dtype=torch.int64) + num_mb = 4 + token_indices = self._indices(num_mb, tokens_per_block, local_bs) + src_block_indices = [1, 3, 0, 2] + src = torch.randn(num_mb, self.num_layers, tokens_per_block, + self.dims, dtype=self.dtype).pin_memory() + batch_scatter_kv_caches( + ptrs, src, token_indices, src_block_indices, + tokens_per_block, self.dims, + block_stride=block_stride, local_block_size=local_bs) + torch.cuda.synchronize() + + caches_cpu = [c.cpu() for c in caches] + for mb in range(num_mb): + for pos in range(tokens_per_block): + flat_idx = token_indices[mb * tokens_per_block + pos] + for layer in range(self.num_layers): + got = _ref_slot(caches_cpu[layer], flat_idx, local_bs) + want = src[src_block_indices[mb], layer, pos, :] + torch.testing.assert_close( + got, want, + msg=f"scatter mismatch mb={mb} pos={pos} " + f"layer={layer} flat={flat_idx}") + if pad: + for layer, c in enumerate(caches_cpu): + self.assertTrue( + bool((c[:, local_bs:, :] == sentinel).all()), + f"scatter wrote into the padding of layer {layer}") + + def test_gather_scatter_roundtrip_strided(self): + """Scattering gathered data into zeroed caches must reproduce exactly + the gathered slots (and only them).""" + local_bs, pad, tokens_per_block = 64, 8, 48 + src_caches = _make_paged_caches( + self.num_layers, self.num_kv_blocks, local_bs, self.dims, + pad, self.device, self.dtype) + dst_caches = _make_paged_caches( + self.num_layers, self.num_kv_blocks, local_bs, self.dims, + pad, self.device, self.dtype, fill_random=False) + block_stride = src_caches[0].stride(0) + src_ptrs = torch.tensor([c.data_ptr() for c in src_caches], + device=self.device, dtype=torch.int64) + dst_ptrs = torch.tensor([c.data_ptr() for c in dst_caches], + device=self.device, dtype=torch.int64) + num_mb = 3 + token_indices = self._indices(num_mb, tokens_per_block, local_bs) + buf = torch.zeros(num_mb, self.num_layers, tokens_per_block, self.dims, + device="cpu", dtype=self.dtype, pin_memory=True) + batch_gather_kv_caches( + src_ptrs, buf, token_indices, list(range(num_mb)), + tokens_per_block, self.dims, + block_stride=block_stride, local_block_size=local_bs) + torch.cuda.synchronize() + batch_scatter_kv_caches( + dst_ptrs, buf, token_indices, list(range(num_mb)), + tokens_per_block, self.dims, + block_stride=block_stride, local_block_size=local_bs) + torch.cuda.synchronize() + src_cpu = [c.cpu() for c in src_caches] + dst_cpu = [c.cpu() for c in dst_caches] + for flat_idx in token_indices: + for layer in range(self.num_layers): + torch.testing.assert_close( + _ref_slot(dst_cpu[layer], flat_idx, local_bs), + _ref_slot(src_cpu[layer], flat_idx, local_bs)) + + +if __name__ == "__main__": + unittest.main() diff --git a/kv_cache_manager/py_connector/test/test_block_translation.py b/kv_cache_manager/py_connector/test/test_block_translation.py new file mode 100644 index 000000000..a75af6484 --- /dev/null +++ b/kv_cache_manager/py_connector/test/test_block_translation.py @@ -0,0 +1,125 @@ +"""Unit tests for the connector's manager-block -> physical-slot translation. + +Covers ``_attn_token_indices`` (attention groups: token-granular three-tier +mapping) and ``_state_block_ids`` (mamba/state groups: manager block's last +token selects the group block), verifying against an independent brute-force +reference implementation, token by token. +""" + +import unittest + +from kv_cache_manager.py_connector.test.vllm_stubs import make_connector +from kv_cache_manager.py_connector.common.types import TransferGroup + + +def _make_group(group_bs, kernel_bs=0, is_attention=True): + return TransferGroup( + group_idx=0, + spec_name="tp0_g0", + is_attention=is_attention, + layer_names=["layer0"], + block_size=group_bs, + per_block_bytes=0, + kernel_block_size=kernel_bs, + ) + + +def _ref_attn_token_indices(manager_bs, group_bs, kernel_bs, manager_block_idxes, + block_table): + """Brute-force reference: walk every token of every manager block and map it + through the block hierarchy step by step.""" + out = [] + for mb in manager_block_idxes: + slots = [] + for tok in range(mb * manager_bs, (mb + 1) * manager_bs): + group_block = tok // group_bs # logical block in group table + tok_in_group = tok - group_block * group_bs + kernel_in_group = tok_in_group // kernel_bs + tok_in_kernel = tok_in_group - kernel_in_group * kernel_bs + physical = block_table[group_block] * (group_bs // kernel_bs) + kernel_in_group + slots.append(physical * kernel_bs + tok_in_kernel) + out.append(slots) + return out + + +def _ref_state_block_ids(manager_bs, group_bs, manager_block_idxes, block_table): + """Brute-force reference: the state covering a manager block is the state of + the group block containing the manager block's last token.""" + out = [] + for mb in manager_block_idxes: + last_token = (mb + 1) * manager_bs - 1 + out.append(block_table[last_token // group_bs]) + return out + + +class TestAttnTokenIndices(unittest.TestCase): + # (manager_bs, group_bs, kernel_bs): ratio=1, ratio>1, manager != group. + CASES = [ + (16, 16, 16), # full attention default: all equal + (32, 16, 16), # preferred_block_size > vllm block size + (528, 528, 64), # hybrid: group block spans several kernel blocks + (528, 528, 528), # hybrid with kernel == group + (48, 16, 8), # manager > group > kernel + ] + + def test_against_reference(self): + for manager_bs, group_bs, kernel_bs in self.CASES: + with self.subTest(manager_bs=manager_bs, group_bs=group_bs, + kernel_bs=kernel_bs): + conn = make_connector(manager_block_size=manager_bs) + group = _make_group(group_bs, kernel_bs) + # Enough non-trivially permuted blocks for 4 manager blocks. + needed = 4 * manager_bs // group_bs + 1 + block_table = [(i * 7 + 3) % 97 for i in range(needed)] + mbis = [0, 1, 3] + got = conn._attn_token_indices(group, mbis, block_table) + want = _ref_attn_token_indices( + manager_bs, group_bs, kernel_bs, mbis, block_table) + self.assertEqual(got, want) + + def test_manual_example(self): + # manager_bs=4, group_bs=2, kernel_bs=2; block_table maps logical + # blocks 0..3 -> physical 5,2,9,0. Manager block 1 covers tokens 4..7 -> + # logical blocks 2,3 -> physical 9,0 -> slots 18,19,0,1. + conn = make_connector(manager_block_size=4) + group = _make_group(group_bs=2, kernel_bs=2) + got = conn._attn_token_indices(group, [1], [5, 2, 9, 0]) + self.assertEqual(got, [[18, 19, 0, 1]]) + + def test_out_of_range_asserts(self): + conn = make_connector(manager_block_size=16) + group = _make_group(group_bs=16, kernel_bs=16) + with self.assertRaises(AssertionError): + conn._attn_token_indices(group, [1], [0]) # table too short + + +class TestStateBlockIds(unittest.TestCase): + def test_against_reference(self): + for manager_bs, group_bs in [(528, 528), (16, 16), (16, 32), (48, 16)]: + with self.subTest(manager_bs=manager_bs, group_bs=group_bs): + conn = make_connector(manager_block_size=manager_bs) + group = _make_group(group_bs, is_attention=False) + needed = 4 * manager_bs // group_bs + 1 + block_table = [(i * 11 + 5) % 89 for i in range(needed)] + mbis = [0, 1, 3] + got = conn._state_block_ids(group, mbis, block_table) + want = _ref_state_block_ids(manager_bs, group_bs, mbis, block_table) + self.assertEqual(got, want) + + def test_manual_example(self): + # manager_bs=4, group_bs=8: manager blocks 0 and 1 both end inside group + # block 0; manager block 2 ends in group block 1. + conn = make_connector(manager_block_size=4) + group = _make_group(group_bs=8, is_attention=False) + got = conn._state_block_ids(group, [0, 1, 2], [7, 3]) + self.assertEqual(got, [7, 7, 3]) + + def test_out_of_range_asserts(self): + conn = make_connector(manager_block_size=16) + group = _make_group(group_bs=16, is_attention=False) + with self.assertRaises(AssertionError): + conn._state_block_ids(group, [2], [0, 1]) + + +if __name__ == "__main__": + unittest.main() diff --git a/kv_cache_manager/py_connector/test/test_data_transfer_results.py b/kv_cache_manager/py_connector/test/test_data_transfer_results.py new file mode 100644 index 000000000..6c1bc5f2e --- /dev/null +++ b/kv_cache_manager/py_connector/test/test_data_transfer_results.py @@ -0,0 +1,136 @@ +"""Unit tests for MultiResult flattening and the save/load done callbacks. + +The done callbacks decode a flat result list whose layout is an implicit +contract with ``_submit_group_tasks``: tasks are submitted group-major +(group0's blocks, then group1's blocks, ...), so a manager block's success is +the stride-AND ``flat[i % num_blocks]``. These tests pin that contract with +hand-computed expectations. +""" + +import threading +import unittest +from unittest.mock import MagicMock + +from kv_cache_manager.py_connector.test import vllm_stubs # noqa: F401 (stubs) +from kv_cache_manager.py_connector.vllm.data_transfer import ( + DataTransferManager, MultiResult) +from kv_cache_manager.py_connector.common.tp_coordinator import ( + CoordinateMsgSerializer) + + +class TestMultiResult(unittest.TestCase): + def test_flatten_in_submission_order(self): + got = [] + mr = MultiResult(3, got.extend) + mr.submit_result(0, [True, False]) + mr.submit_result(1, [False]) + mr.submit_result(2, [True, True, True]) + self.assertEqual(got, [True, False, False, True, True, True]) + + def test_out_of_order_submit(self): + got = [] + mr = MultiResult(3, got.extend) + mr.submit_result(2, ["c"]) + mr.submit_result(0, ["a"]) + self.assertEqual(got, []) # callback must not fire early + mr.submit_result(1, ["b"]) + self.assertEqual(got, ["a", "b", "c"]) + + def test_duplicate_submit_asserts(self): + mr = MultiResult(2, lambda flat: None) + mr.submit_result(0, [True]) + with self.assertRaises(AssertionError): + mr.submit_result(0, [True]) + + def test_concurrent_submit(self): + n = 64 + results = [] + done = threading.Event() + + def cb(flat): + results.append(flat) + done.set() + + mr = MultiResult(n, cb) + barrier = threading.Barrier(n) + + def worker(i): + barrier.wait() + mr.submit_result(i, [i]) + + threads = [threading.Thread(target=worker, args=(i,)) for i in range(n)] + for t in threads: + t.start() + for t in threads: + t.join() + self.assertTrue(done.wait(timeout=5)) + self.assertEqual(len(results), 1) # callback fires exactly once + self.assertEqual(results[0], list(range(n))) + + +def _make_dtm(): + """DataTransferManager with only the state the callbacks touch.""" + dtm = DataTransferManager.__new__(DataTransferManager) + dtm._coordinator_client = MagicMock() + return dtm + + +def _sent_event(dtm): + (payload,), _ = dtm._coordinator_client.send.call_args + return CoordinateMsgSerializer.loads(payload).content + + +class TestSaveDoneCallback(unittest.TestCase): + def test_multi_group_stride_and(self): + # 3 blocks x 2 groups, flat = group0[b0,b1,b2] + group1[b0,b1,b2]. + # Block b is saved only if both groups succeeded for b. + dtm = _make_dtm() + cb = dtm.create_save_done_callback("req", 0, "sess", num_blocks=3) + cb([True, True, False, # group 0 + True, False, True]) # group 1 + evt = _sent_event(dtm) + self.assertEqual(evt.type, "SendBlockFinishedEvent") + self.assertEqual(evt.write_session_id, "sess") + self.assertEqual(evt.is_success_list, [True, False, False]) + + def test_single_group_passthrough(self): + dtm = _make_dtm() + cb = dtm.create_save_done_callback("req", 1, "sess", num_blocks=2) + cb([False, True]) + self.assertEqual(_sent_event(dtm).is_success_list, [False, True]) + + +class TestLoadDoneCallback(unittest.TestCase): + def test_multi_group_failure_merge(self): + dtm = _make_dtm() + cb = dtm.create_load_done_callback( + "req", 0, epoch=7, block_ids=[10, 20, 30], num_blocks=3) + cb([True, False, True, # group 0 + True, True, False]) # group 1 + evt = _sent_event(dtm) + self.assertEqual(evt.type, "LoadBlockFinishedEvent") + self.assertEqual(evt.epoch, 7) + # blocks 1 and 2 each failed in one group -> report their table ids. + self.assertEqual(evt.failed_block_idxs, [20, 30]) + + def test_all_success_reports_empty(self): + dtm = _make_dtm() + cb = dtm.create_load_done_callback( + "req", 0, epoch=0, block_ids=[10, 20], num_blocks=2) + cb([True, True, True, True]) + self.assertEqual(_sent_event(dtm).failed_block_idxs, []) + + def test_report_failures_false_hybrid(self): + # Hybrid models cannot report invalid block ids to vLLM: the failure + # must be swallowed (empty failed list) but the finished event still sent. + dtm = _make_dtm() + cb = dtm.create_load_done_callback( + "req", 0, epoch=1, block_ids=[], num_blocks=2, report_failures=False) + cb([False, True]) + evt = _sent_event(dtm) + self.assertEqual(evt.type, "LoadBlockFinishedEvent") + self.assertEqual(evt.failed_block_idxs, []) + + +if __name__ == "__main__": + unittest.main() diff --git a/kv_cache_manager/py_connector/test/test_scheduler_state.py b/kv_cache_manager/py_connector/test/test_scheduler_state.py new file mode 100644 index 000000000..8b74177d4 --- /dev/null +++ b/kv_cache_manager/py_connector/test/test_scheduler_state.py @@ -0,0 +1,391 @@ +"""Unit tests for the connector's scheduler-side logic. + +Covers, against fake vLLM SchedulerOutput / Request objects: + +* ``get_num_new_matched_tokens`` -- including the full-prompt external hit cap + (a fully cached prompt must leave >= 1 token to recompute, otherwise vLLM's + synchronous-load scheduling path asserts ``num_new_tokens > 0``); +* ``parse_block_mask_to_save_indices`` -- ``offset`` and ``bool_masks`` forms; +* ``_parse_groups`` -- full-attention single group, hybrid multi group, eagle + group skip, unsupported spec error; +* ``build_connector_meta`` -- new request, cached deltas (``new_block_ids`` + None / non-None), preemption resume via both the 0.26 ``resumed_req_ids`` and + the legacy ``resumed_from_preemption`` interfaces, save-threshold trigger, + and the two ``request_finished`` paths (saves landed / in-flight). +""" + +import unittest +from dataclasses import dataclass, field +from types import SimpleNamespace +from unittest.mock import MagicMock + +from kv_cache_manager.py_connector.test.vllm_stubs import ( + make_connector, ReqState, GroupMeta) +from kv_cache_manager.py_connector.vllm.v1_connector import TairKvCacheConnector +from kv_cache_manager.py_connector.vllm.metadata import SaveRequest + + +# --------------------------------------------------------------------------- # +# Fakes +# --------------------------------------------------------------------------- # +@dataclass +class FakeRequest: + request_id: str + prompt_token_ids: list + output_token_ids: list = field(default_factory=list) + + @property + def num_tokens(self): + return len(self.prompt_token_ids) + len(self.output_token_ids) + + @property + def all_token_ids(self): + return self.prompt_token_ids + self.output_token_ids + + +def make_scheduler_connector(mbs=16, vllm_bs=None, locations=None): + """Connector with the scheduler-side state build_connector_meta needs.""" + conn = make_connector(manager_block_size=mbs, vllm_block_size=vllm_bs) + conn._epoch = 0 + conn._alive_requests = {} + conn._waiting_to_load_requests = [] + import threading + conn._waiting_to_save_requests_lock = threading.Lock() + conn._waiting_to_save_requests = [] + conn._waiting_to_finish_requests = [] + conn._canceled_save_request_ids_lock = threading.Lock() + conn._canceled_save_request_ids = [] + conn._http_executor = MagicMock() + conn._location_query_manager = MagicMock() + conn._location_query_manager.get_locations_for_query.return_value = ( + True, locations if locations is not None else []) + return conn + + +def fake_scheduler_output(new_reqs=(), cached_req_ids=(), num_scheduled=None, + new_block_ids=(), resumed_req_ids=frozenset(), + legacy_resumed=None): + """Build a fake SchedulerOutput. legacy_resumed switches the cached-reqs + container to the pre-0.26 interface (resumed_from_preemption list, no + resumed_req_ids attribute).""" + if legacy_resumed is not None: + cached = SimpleNamespace( + req_ids=list(cached_req_ids), + resumed_from_preemption=list(legacy_resumed), + new_block_ids=list(new_block_ids), + ) + else: + cached = SimpleNamespace( + req_ids=list(cached_req_ids), + resumed_req_ids=set(resumed_req_ids), + new_block_ids=list(new_block_ids), + ) + return SimpleNamespace( + scheduled_new_reqs=list(new_reqs), + scheduled_cached_reqs=cached, + num_scheduled_tokens=dict(num_scheduled or {}), + ) + + +def make_locations(n): + return [{"location_specs": [{"name": "tp0_g0", "uri": f"file://blk{i}"}]} + for i in range(n)] + + +# --------------------------------------------------------------------------- # +# get_num_new_matched_tokens +# --------------------------------------------------------------------------- # +class TestGetNumNewMatchedTokens(unittest.TestCase): + MBS = 16 + + def _run(self, prompt_len, num_computed, num_locations): + conn = make_scheduler_connector( + mbs=self.MBS, locations=make_locations(num_locations)) + req = FakeRequest("r0", list(range(prompt_len))) + matched, async_load = conn.get_num_new_matched_tokens(req, num_computed) + return conn, matched, async_load + + def test_partial_hit_no_cap(self): + conn, matched, async_load = self._run(4 * self.MBS + 5, 0, 4) + self.assertEqual(matched, 4 * self.MBS) + self.assertTrue(async_load) # a pending load is reported as async + self.assertEqual(conn._waiting_to_load_requests[0].manager_block_idxes, + [0, 1, 2, 3]) + + def test_full_hit_capped_to_leave_one_token(self): + # Prompt is exactly 4 manager blocks, all externally cached: the last + # block must be dropped so vLLM still schedules >= 1 new token. + conn, matched, _ = self._run(4 * self.MBS, 0, 4) + self.assertEqual(matched, 3 * self.MBS) + self.assertEqual(conn._waiting_to_load_requests[0].manager_block_idxes, + [0, 1, 2]) + # has_saved_block_num counts only the blocks actually treated as hit. + self.assertEqual(conn._alive_requests["r0"].has_saved_block_num, 3) + + def test_full_hit_with_local_prefix(self): + # 2 blocks locally computed + 2 remote = whole prompt -> drop one remote. + conn, matched, _ = self._run(4 * self.MBS, 2 * self.MBS, 2) + self.assertEqual(matched, self.MBS) + self.assertEqual(conn._waiting_to_load_requests[0].manager_block_idxes, [2]) + + def test_single_block_full_hit_degrades_to_zero(self): + conn, matched, async_load = self._run(self.MBS, 0, 1) + self.assertEqual(matched, 0) + self.assertFalse(async_load) + self.assertEqual(conn._waiting_to_load_requests, []) + + def test_no_locations(self): + conn, matched, async_load = self._run(100, 0, 0) + self.assertEqual(matched, 0) + self.assertFalse(async_load) + + +# --------------------------------------------------------------------------- # +# parse_block_mask_to_save_indices +# --------------------------------------------------------------------------- # +class TestParseBlockMask(unittest.TestCase): + def setUp(self): + self.conn = make_connector() + + def test_offset_branch(self): + resp = {"block_mask": {"offset": 2}} + self.assertEqual( + self.conn.parse_block_mask_to_save_indices(resp, 5), [2, 3, 4]) + + def test_offset_zero(self): + resp = {"block_mask": {"offset": 0}} + self.assertEqual( + self.conn.parse_block_mask_to_save_indices(resp, 3), [0, 1, 2]) + + def test_bool_masks_branch(self): + resp = {"block_mask": {"bool_masks": {"values": [True, False, True, False]}}} + self.assertEqual( + self.conn.parse_block_mask_to_save_indices(resp, 4), [1, 3]) + + def test_missing_mask(self): + self.assertEqual(self.conn.parse_block_mask_to_save_indices({}, 3), []) + + +# --------------------------------------------------------------------------- # +# _parse_groups +# --------------------------------------------------------------------------- # +class TestParseGroups(unittest.TestCase): + def _kv_cache_config(self, groups): + return SimpleNamespace(kv_cache_groups=groups) + + def _attn_group(self, layers, block_size=16, page_size_bytes=32768): + from vllm.v1.kv_cache_interface import FullAttentionSpec + return SimpleNamespace( + layer_names=layers, + kv_cache_spec=FullAttentionSpec(block_size, page_size_bytes)) + + def _mamba_group(self, layers, block_size=528, page_size_bytes=1024): + from vllm.v1.kv_cache_interface import MambaSpec + return SimpleNamespace( + layer_names=layers, + kv_cache_spec=MambaSpec(block_size, page_size_bytes)) + + def test_full_attention_single_group(self): + conn = make_connector(manager_block_size=32) + metas = conn._parse_groups(self._kv_cache_config( + [self._attn_group(["l0", "l1"], block_size=16, page_size_bytes=32768)])) + self.assertEqual(len(metas), 1) + m = metas[0] + self.assertTrue(m.is_attention) + self.assertEqual(m.group_idx, 0) + self.assertEqual(m.block_size, 16) + # per_token = 32768 // 16 = 2048; per_block = 2048 * 32 (manager) * 2 layers + self.assertEqual(m.per_block_bytes, 2048 * 32 * 2) + + def test_hybrid_multi_group(self): + conn = make_connector(manager_block_size=528) + metas = conn._parse_groups(self._kv_cache_config([ + self._mamba_group(["m0", "m1"], page_size_bytes=1000), + self._mamba_group(["m2"], page_size_bytes=2000), + self._attn_group(["a0"], block_size=528, page_size_bytes=528 * 64), + ])) + self.assertEqual([m.group_idx for m in metas], [0, 1, 2]) + self.assertEqual([m.is_attention for m in metas], [False, False, True]) + self.assertEqual(metas[0].per_block_bytes, 1000 * 2) # page * layers + self.assertEqual(metas[1].per_block_bytes, 2000) + self.assertEqual(metas[2].per_block_bytes, 64 * 528) # per_token * mbs + + def test_eagle_group_skipped(self): + conn = make_connector() + eagle = self._attn_group(["drafter"]) + eagle.is_eagle_group = True + metas = conn._parse_groups(self._kv_cache_config( + [eagle, self._attn_group(["a0"])])) + self.assertEqual(len(metas), 1) + self.assertEqual(metas[0].layer_names, ["a0"]) + self.assertEqual(metas[0].group_idx, 1) # group_idx keeps vLLM numbering + + def test_unsupported_spec_raises(self): + conn = make_connector() + bad = SimpleNamespace(layer_names=["x"], kv_cache_spec=object()) + with self.assertRaises(NotImplementedError): + conn._parse_groups(self._kv_cache_config([bad])) + + def test_no_usable_groups_asserts(self): + conn = make_connector() + with self.assertRaises(AssertionError): + conn._parse_groups(self._kv_cache_config([])) + + +# --------------------------------------------------------------------------- # +# build_connector_meta +# --------------------------------------------------------------------------- # +class TestBuildConnectorMeta(unittest.TestCase): + MBS = 16 + + def _new_request(self, conn, req_id, num_tokens, num_blocks, + num_locations=0): + """Simulate the scheduler flow for a fresh request: query, alloc, then + one build_connector_meta step.""" + conn._location_query_manager.get_locations_for_query.return_value = ( + True, make_locations(num_locations)) + req = FakeRequest(req_id, list(range(num_tokens))) + conn.get_num_new_matched_tokens(req, 0) + block_ids = [list(range(100, 100 + num_blocks))] + conn.update_state_after_alloc( + req, SimpleNamespace(get_block_ids=lambda: block_ids), 0) + out = fake_scheduler_output( + new_reqs=[SimpleNamespace(req_id=req_id, block_ids=block_ids)]) + return req, conn.build_connector_meta(out) + + def test_new_request_full_state(self): + conn = make_scheduler_connector(mbs=self.MBS) + req, meta = self._new_request(conn, "r0", 40, 3) + self.assertEqual(len(meta.requests), 1) + delta = meta.requests[0] + self.assertFalse(delta.is_delta) + self.assertEqual(delta.new_tokens_ids, list(range(40))) + self.assertEqual(delta.new_block_ids_per_group, [[100, 101, 102]]) + # 40 tokens / 3 blocks -> min(40, 48)//16 = 2 blocks to save. + conn._http_executor.submit.assert_called_once() + args = conn._http_executor.submit.call_args[0] + self.assertEqual(args[1:], ("r0", list(range(32)), 2)) + self.assertEqual(conn._alive_requests["r0"].has_saved_block_num, 2) + + def test_load_request_emitted_after_alloc(self): + conn = make_scheduler_connector(mbs=self.MBS) + req, meta = self._new_request(conn, "r0", 40, 3, num_locations=2) + self.assertEqual(len(meta.to_load_requests), 1) + lr = meta.to_load_requests[0] + self.assertEqual(lr.manager_block_idxes, [0, 1]) + self.assertEqual(lr.all_block_ids, [[100, 101, 102]]) + # Externally hit blocks are not re-saved. + conn._http_executor.submit.assert_not_called() + + def test_cached_delta_with_and_without_new_blocks(self): + conn = make_scheduler_connector(mbs=self.MBS) + req, _ = self._new_request(conn, "r0", 40, 3) + # Step 2: 8 decode tokens, no new blocks (PR #23262: may be None). + req.output_token_ids = list(range(1000, 1008)) + out = fake_scheduler_output( + cached_req_ids=["r0"], num_scheduled={"r0": 8}, new_block_ids=[None]) + meta = conn.build_connector_meta(out) + delta = meta.requests[0] + self.assertTrue(delta.is_delta) + self.assertEqual(delta.new_tokens_ids, list(range(1000, 1008))) + self.assertEqual(delta.new_block_ids_per_group, []) + # Step 3: 2 more tokens with a new block -> table grows. + req.output_token_ids = list(range(1000, 1010)) + out = fake_scheduler_output( + cached_req_ids=["r0"], num_scheduled={"r0": 2}, + new_block_ids=[[[103]]]) + meta = conn.build_connector_meta(out) + self.assertEqual(meta.requests[0].new_block_ids_per_group, [[103]]) + self.assertEqual(conn._alive_requests["r0"].block_ids_per_group, + [[100, 101, 102, 103]]) + + def _preempted_step(self, conn, req, use_legacy): + kwargs = dict(cached_req_ids=["r0"], num_scheduled={"r0": 0}, + new_block_ids=[[[200, 201]]]) + if use_legacy: + kwargs["legacy_resumed"] = [True] + else: + kwargs["resumed_req_ids"] = {"r0"} + return conn.build_connector_meta(fake_scheduler_output(**kwargs)) + + def test_resumed_from_preemption_both_interfaces(self): + for use_legacy in (False, True): + with self.subTest(legacy=use_legacy): + conn = make_scheduler_connector(mbs=self.MBS) + req, _ = self._new_request(conn, "r0", 40, 3) + meta = self._preempted_step(conn, req, use_legacy) + delta = meta.requests[0] + self.assertTrue(delta.resumed_from_preemption) + # Resume replaces (not extends) the block table. + self.assertEqual(conn._alive_requests["r0"].block_ids_per_group, + [[200, 201]]) + + def test_save_threshold_grows_incrementally(self): + conn = make_scheduler_connector(mbs=self.MBS) + req, _ = self._new_request(conn, "r0", 40, 3) # saved 2 blocks + conn._http_executor.submit.reset_mock() + # 8 more tokens -> 48 total, table full at 3 blocks -> third block saves. + req.output_token_ids = list(range(1000, 1008)) + out = fake_scheduler_output( + cached_req_ids=["r0"], num_scheduled={"r0": 8}, new_block_ids=[[[103]]]) + conn.build_connector_meta(out) + args = conn._http_executor.submit.call_args[0] + self.assertEqual(args[3], 3) # target_save_num + self.assertEqual(conn._alive_requests["r0"].has_saved_block_num, 3) + + def test_save_request_drain_and_finish_paths(self): + conn = make_scheduler_connector(mbs=self.MBS) + req, _ = self._new_request(conn, "r0", 40, 3) + state = conn._alive_requests["r0"] + self.assertEqual(state.scheduled_saving_count, 1) + + # Finish while the save is still in flight: request must stay alive. + keep, extra = conn.request_finished(req, []) + self.assertTrue(keep) + self.assertTrue(state.need_report_after_saving_finished) + self.assertIn("r0", conn._alive_requests) + + # The async save lands: drained into to_save_requests and, because the + # request already finished, a FinishRequest is emitted and state dropped. + with conn._waiting_to_save_requests_lock: + conn._waiting_to_save_requests.append( + SaveRequest("r0", make_locations(2), [0, 1], "sess")) + meta = conn.build_connector_meta(fake_scheduler_output()) + self.assertEqual(len(meta.to_save_requests), 1) + self.assertEqual([f.req_id for f in meta.to_finish_requests], ["r0"]) + self.assertNotIn("r0", conn._alive_requests) + + def test_request_finished_when_saves_landed(self): + conn = make_scheduler_connector(mbs=self.MBS) + req, _ = self._new_request(conn, "r0", 40, 3) + with conn._waiting_to_save_requests_lock: + conn._waiting_to_save_requests.append( + SaveRequest("r0", make_locations(2), [0, 1], "sess")) + conn.build_connector_meta(fake_scheduler_output()) + keep, extra = conn.request_finished(req, []) + self.assertTrue(keep) + self.assertNotIn("r0", conn._alive_requests) + meta = conn.build_connector_meta(fake_scheduler_output()) + self.assertEqual([f.req_id for f in meta.to_finish_requests], ["r0"]) + + def test_canceled_save_unknown_request_no_crash(self): + # Cancellations arrive from http_executor threads and may race request + # teardown; an unknown req_id must be skipped, not KeyError. + conn = make_scheduler_connector(mbs=self.MBS) + with conn._canceled_save_request_ids_lock: + conn._canceled_save_request_ids.append("ghost") + conn.build_connector_meta(fake_scheduler_output()) # must not raise + + def test_canceled_save_finishes_request(self): + conn = make_scheduler_connector(mbs=self.MBS) + req, _ = self._new_request(conn, "r0", 40, 3) + conn.request_finished(req, []) # save in flight -> delayed finish + with conn._canceled_save_request_ids_lock: + conn._canceled_save_request_ids.append("r0") + meta = conn.build_connector_meta(fake_scheduler_output()) + self.assertEqual([f.req_id for f in meta.to_finish_requests], ["r0"]) + self.assertNotIn("r0", conn._alive_requests) + + +if __name__ == "__main__": + unittest.main() diff --git a/kv_cache_manager/py_connector/test/vllm_stubs.py b/kv_cache_manager/py_connector/test/vllm_stubs.py new file mode 100644 index 000000000..8c6644c17 --- /dev/null +++ b/kv_cache_manager/py_connector/test/vllm_stubs.py @@ -0,0 +1,133 @@ +"""Shared test stubs: make ``v1_connector`` importable without vLLM/CUDA/pybind. + +``v1_connector`` imports vLLM and the compiled ``kvcm_py_client`` at module +level. For pure-logic unit tests we register lightweight stand-ins in +``sys.modules`` *before* the first import, then build connector instances via +``__new__`` with only the attributes the code under test reads. No production +module is modified. +""" + +import sys +import types +from typing import Optional +from unittest.mock import MagicMock + + +def _module(name: str) -> types.ModuleType: + mod = sys.modules.get(name) + if mod is None: + mod = types.ModuleType(name) + sys.modules[name] = mod + return mod + + +def _install_stubs(): + existing = sys.modules.get("vllm") + if existing is not None: + # Either our stub is already in place or the real vLLM is importable; + # in both cases the connector import will succeed as-is. + return + + # ---- kv_cache_manager.client.pybind (compiled extension) ---- + pybind = _module("kv_cache_manager.client.pybind") + kvcm_py_client = MagicMock() + kvcm_py_client.ClientErrorCode.ER_OK = 0 + pybind.kvcm_py_client = kvcm_py_client + + # ---- kv_cache_manager.py_connector.common._version_info (generated) ---- + version = _module("kv_cache_manager.py_connector.common._version_info") + version.FULL_VERSION = "0.0.0-test" + version.GIT_COMMIT = "test" + version.BUILD_TIME = "test" + + # ---- vllm ---- + vllm = _module("vllm") + vllm._kvcm_test_stub = True + + config = _module("vllm.config") + config.VllmConfig = MagicMock + vllm.config = config + + distributed = _module("vllm.distributed") + distributed.get_tensor_model_parallel_rank = lambda: 0 + vllm.distributed = distributed + _module("vllm.distributed.kv_transfer") + _module("vllm.distributed.kv_transfer.kv_connector") + _module("vllm.distributed.kv_transfer.kv_connector.v1") + base = _module("vllm.distributed.kv_transfer.kv_connector.v1.base") + + class KVConnectorRole: + SCHEDULER = 0 + WORKER = 1 + + class KVConnectorMetadata: + pass + + class KVConnectorBase_V1: + def __init__(self, vllm_config, role, kv_cache_config=None): + self._connector_metadata = None + + def _get_connector_metadata(self): + return self._connector_metadata + + class SupportsHMA: + pass + + base.KVConnectorBase_V1 = KVConnectorBase_V1 + base.KVConnectorMetadata = KVConnectorMetadata + base.KVConnectorRole = KVConnectorRole + base.SupportsHMA = SupportsHMA + + utils = _module("vllm.utils") + torch_utils = _module("vllm.utils.torch_utils") + torch_utils.get_kv_cache_torch_dtype = MagicMock() + network_utils = _module("vllm.utils.network_utils") + network_utils.get_ip = lambda: "127.0.0.1" + utils.torch_utils = torch_utils + utils.network_utils = network_utils + + v1 = _module("vllm.v1") + kv_cache_interface = _module("vllm.v1.kv_cache_interface") + + class FullAttentionSpec: + def __init__(self, block_size, page_size_bytes): + self.block_size = block_size + self.page_size_bytes = page_size_bytes + + class MambaSpec: + def __init__(self, block_size, page_size_bytes): + self.block_size = block_size + self.page_size_bytes = page_size_bytes + + kv_cache_interface.FullAttentionSpec = FullAttentionSpec + kv_cache_interface.MambaSpec = MambaSpec + + _module("vllm.v1.core") + sched = _module("vllm.v1.core.sched") + output = _module("vllm.v1.core.sched.output") + output.SchedulerOutput = MagicMock + sched.output = output + + outputs = _module("vllm.v1.outputs") + outputs.KVConnectorOutput = MagicMock + v1.kv_cache_interface = kv_cache_interface + v1.outputs = outputs + + +_install_stubs() + +# Import after stubs are in place. +from kv_cache_manager.py_connector.vllm.v1_connector import ( # noqa: E402 + TairKvCacheConnector, GroupMeta, ReqState) + + +def make_connector(manager_block_size: int = 16, + vllm_block_size: Optional[int] = None, + num_groups: int = 1) -> TairKvCacheConnector: + """Build a bare TairKvCacheConnector (no __init__) with the minimal state + used by the pure translation / scheduler-side logic under test.""" + conn = TairKvCacheConnector.__new__(TairKvCacheConnector) + conn._manager_block_size = manager_block_size + conn._vllm_block_size = vllm_block_size or manager_block_size + conn._num_groups = num_groups + return conn diff --git a/kv_cache_manager/py_connector/vllm/BUILD b/kv_cache_manager/py_connector/vllm/BUILD index 2fc64b26c..acf24672e 100644 --- a/kv_cache_manager/py_connector/vllm/BUILD +++ b/kv_cache_manager/py_connector/vllm/BUILD @@ -7,6 +7,7 @@ load("@python_platform//:platform.bzl", "python_platform") py_library( name = "vllm_connector", srcs = glob(["*.py"]), + visibility = ["//kv_cache_manager/py_connector:__subpackages__"], deps = ["//kv_cache_manager/py_connector/common:common", "//kv_cache_manager/py_connector/kernel:kernel", "//kv_cache_manager/client/pybind:kvcm_py_client_lib"] From 98ef717bffef719c77d2e8fa315ba3a7dfb0fef7 Mon Sep 17 00:00:00 2001 From: xiaozeyu Date: Tue, 28 Jul 2026 18:18:08 +0800 Subject: [PATCH 5/7] [py_connector] fix load-failure retry loop and null mamba state transfer Two production bugs surfaced by the new e2e scenarios: 1. Fail-reschedule loop after a KV load failure. With kv_load_failure_policy=recompute, vLLM reschedules the request and calls get_num_new_matched_tokens again; the manager still advertises the blocks whose storage is gone, so the connector re-matched them and the engine looped load-fail-reschedule forever (request hung). A request that already went through an external load attempt (blocks were allocated) now skips external matching on re-query and recomputes locally. 2. Multi-block hybrid saves always failed. vLLM's mamba 'align' mode only materializes the state block ending the matched region (single_type_kv_cache_manager.MambaManager assigns the null block to interior positions), so every interior manager block has a null (id 0) state target by design. save_task treated that as a failure, the stride-AND merge then dropped those manager blocks from the manager's prefix chain, and hybrid caching silently degraded to the final block only. Null state targets are now transferred vacuously (reported success, nothing copied) on both save and load; load_task previously failed the whole task on any null target for the same reason. Covered by test_scheduler_state (retry guard), test_data_transfer_results (vacuous null-state save/load), and e2e: hybrid basic now verifies 4/4 manager blocks bit-exact; test_load_failure exercises the recompute path end to end. --- .../test/test_data_transfer_results.py | 37 +++++++++++ .../py_connector/test/test_scheduler_state.py | 22 +++++++ .../py_connector/vllm/data_transfer.py | 65 +++++++++++++------ .../py_connector/vllm/v1_connector.py | 16 +++++ 4 files changed, 121 insertions(+), 19 deletions(-) diff --git a/kv_cache_manager/py_connector/test/test_data_transfer_results.py b/kv_cache_manager/py_connector/test/test_data_transfer_results.py index 6c1bc5f2e..7e43e9f7e 100644 --- a/kv_cache_manager/py_connector/test/test_data_transfer_results.py +++ b/kv_cache_manager/py_connector/test/test_data_transfer_results.py @@ -132,5 +132,42 @@ def test_report_failures_false_hybrid(self): self.assertEqual(evt.failed_block_idxs, []) +class TestNullStateBlocks(unittest.TestCase): + """Mamba 'align' mode: null (id 0) state targets carry no state by design + (vLLM materializes states only at segment boundaries), so save/load must + treat them as vacuous successes -- failing them would stride-AND whole + manager blocks out of the manager's prefix chain and kill multi-block + hybrid caching. The all-null path takes no GPU work, so it runs on CPU.""" + + @staticmethod + def _state_group(): + from kv_cache_manager.py_connector.common.types import TransferGroup + return TransferGroup( + group_idx=0, spec_name="tp0_g0", is_attention=False, + layer_names=["m0"], block_size=528, per_block_bytes=1024, + kernel_block_size=528) + + def test_save_all_null_blocks_vacuously_succeed(self): + dtm = _make_dtm() + results = {} + mr = MultiResult(1, lambda flat: results.setdefault("flat", flat)) + dtm.save_task(mr, 0, self._state_group(), + remote_uris=["u0", "u1"], + block_token_indices=None, + block_ids=[0, 0], + ready_event=None) + self.assertEqual(results["flat"], [True, True]) + + def test_load_all_null_blocks_vacuously_succeed(self): + dtm = _make_dtm() + results = {} + mr = MultiResult(1, lambda flat: results.setdefault("flat", flat)) + dtm.load_task(mr, 0, self._state_group(), + remote_uris=["u0", "u1", "u2"], + block_token_indices=None, + block_ids=[0, 0, 0]) + self.assertEqual(results["flat"], [True, True, True]) + + if __name__ == "__main__": unittest.main() diff --git a/kv_cache_manager/py_connector/test/test_scheduler_state.py b/kv_cache_manager/py_connector/test/test_scheduler_state.py index 8b74177d4..c4f70ead6 100644 --- a/kv_cache_manager/py_connector/test/test_scheduler_state.py +++ b/kv_cache_manager/py_connector/test/test_scheduler_state.py @@ -139,6 +139,28 @@ def test_no_locations(self): self.assertEqual(matched, 0) self.assertFalse(async_load) + def test_requery_after_load_attempt_skips_external(self): + # A request that already went through an external load (blocks were + # allocated) and returned to WAITING -- KV load failure with + # policy=recompute, or preemption -- must not re-match: the manager may + # still advertise blocks whose storage is gone, and re-matching loops + # fail -> reschedule forever. + conn = make_scheduler_connector(mbs=self.MBS, locations=make_locations(2)) + req = FakeRequest("r0", list(range(4 * self.MBS + 5))) + matched, _ = conn.get_num_new_matched_tokens(req, 0) + self.assertEqual(matched, 2 * self.MBS) + # vLLM allocates blocks for the load attempt. + conn.update_state_after_alloc( + req, SimpleNamespace(get_block_ids=lambda: [[100, 101, 102]]), matched) + # Retry: same request re-enters the waiting queue with 0 computed. + matched2, async2 = conn.get_num_new_matched_tokens(req, 0) + self.assertEqual(matched2, 0) + self.assertFalse(async2) + self.assertEqual(len(conn._waiting_to_load_requests), 1) # no new load + state = conn._alive_requests["r0"] + self.assertEqual(state.remote_matched_token_num, 0) + self.assertEqual(state.has_saved_block_num, 0) + # --------------------------------------------------------------------------- # # parse_block_mask_to_save_indices diff --git a/kv_cache_manager/py_connector/vllm/data_transfer.py b/kv_cache_manager/py_connector/vllm/data_transfer.py index eba0d1e63..0642d1fe3 100644 --- a/kv_cache_manager/py_connector/vllm/data_transfer.py +++ b/kv_cache_manager/py_connector/vllm/data_transfer.py @@ -113,18 +113,31 @@ def save_task(self, multi_result: MultiResult, task_idx, group: TransferGroup, block_token_indices: attention -> list[list[int]] flat token slots per block. block_ids: state -> list[int] block id per manager block; - id 0 is vLLM's null block: the boundary state was - never materialized, so that block cannot be saved. + id 0 is vLLM's null block: no state exists at that + boundary by design (mamba "align" sparse states), + the block is reported saved vacuously. """ n = len(remote_uris) if group.is_attention: valid = list(range(n)) else: + # vLLM's mamba "align" mode only materializes the state block at + # segment boundaries (single_type_kv_cache_manager.MambaManager + # allocates the null block for intermediate positions; a hit only + # ever consumes the state ending the matched region). A null + # (id 0) state target therefore means "no state exists by design", + # not a failure: report it saved vacuously. Failing it would + # stride-AND the whole manager block out of the manager's prefix + # chain and kill multi-block hybrid caching entirely. valid = [i for i in range(n) if block_ids[i] != 0] if len(valid) < n: - logger.warning("save group %s: %d/%d blocks have no materialized " - "state, failing them", group.spec_name, n - len(valid), n) - ok_mask = [False] * n + logger.info("save group %s: %d/%d blocks have no materialized " + "state, saving vacuously", group.spec_name, + n - len(valid), n) + # Vacuous (skipped) blocks succeed; transferred blocks start False and + # are flipped by the transfer result below. + valid_set = set(valid) + ok_mask = [i not in valid_set for i in range(n)] if valid: cpu_buffer = torch.empty(len(valid) * group.per_block_bytes, dtype=torch.uint8, device="cpu", pin_memory=True) @@ -184,17 +197,29 @@ def cb(flat): def load_task(self, multi_result: MultiResult, task_idx, group: TransferGroup, remote_uris, block_token_indices, block_ids): n = len(remote_uris) - if not group.is_attention and any(b == 0 for b in block_ids): - # Null block: nowhere to scatter the state. Should not happen for - # loads (vLLM allocates real blocks for external tokens). - logger.warning("load group %s: null block in targets, failing task", - group.spec_name) - multi_result.submit_result(task_idx, [False] * n) + if group.is_attention: + valid = list(range(n)) + else: + # Mirror of save_task: in mamba "align" mode vLLM only allocates a + # real state block for the final matched boundary; intermediate + # manager blocks get the null block (their state is not needed to + # resume). Skip them vacuously and load only materialized targets. + valid = [i for i in range(n) if block_ids[i] != 0] + if len(valid) < n: + logger.info("load group %s: %d/%d blocks have null state " + "targets, skipping them", group.spec_name, + n - len(valid), n) + valid_set = set(valid) + ok_mask = [i not in valid_set for i in range(n)] + if not valid: + multi_result.submit_result(task_idx, ok_mask) return - cpu_buffer = torch.empty(n * group.per_block_bytes, dtype=torch.uint8, + cpu_buffer = torch.empty(len(valid) * group.per_block_bytes, dtype=torch.uint8, device="cpu", pin_memory=True) - buffers = self._make_block_buffers(cpu_buffer.data_ptr(), group.per_block_bytes, n) - result = self._transfer_client.LoadKvCaches(remote_uris, buffers) + buffers = self._make_block_buffers(cpu_buffer.data_ptr(), + group.per_block_bytes, len(valid)) + uris = [remote_uris[i] for i in valid] + result = self._transfer_client.LoadKvCaches(uris, buffers) ok = (result == kvcm_py_client.ClientErrorCode.ER_OK) if ok: with self._device_mod.stream(self._load_stream): @@ -208,18 +233,20 @@ def load_task(self, multi_result: MultiResult, task_idx, group: TransferGroup, kv_stride=group.kv_stride, block_stride=group.block_stride, local_block_size=group.kernel_block_size) else: - for i, block_id in enumerate(block_ids): + for out_i, i in enumerate(valid): for layer_idx in range(group.layer_num): - src = (i * group.layer_num + layer_idx) * group.page_size_bytes - group.block_view_tensors[layer_idx][block_id].copy_( + src = (out_i * group.layer_num + layer_idx) * group.page_size_bytes + group.block_view_tensors[layer_idx][block_ids[i]].copy_( gpu_buffer[src:src + group.page_size_bytes]) done = self._device_mod.Event() done.record(self._load_stream) done.synchronize() else: logger.warning("load task failed group=%s uris=%d result=%s", - group.spec_name, n, result) - multi_result.submit_result(task_idx, [ok] * n) + group.spec_name, len(uris), result) + for i in valid: + ok_mask[i] = ok + multi_result.submit_result(task_idx, ok_mask) def create_load_done_callback(self, req_id, tp_rank, epoch, block_ids, num_blocks, report_failures=True): diff --git a/kv_cache_manager/py_connector/vllm/v1_connector.py b/kv_cache_manager/py_connector/vllm/v1_connector.py index baa8b706a..8eeb4eb5f 100644 --- a/kv_cache_manager/py_connector/vllm/v1_connector.py +++ b/kv_cache_manager/py_connector/vllm/v1_connector.py @@ -700,6 +700,22 @@ def bind_connector_metadata(self, connector_metadata: KVConnectorMetadata) -> No # ------------------------------------------------------------------ # def get_num_new_matched_tokens(self, request: "Request", num_computed_tokens: int) -> Tuple[Optional[int], bool]: + prev = self._alive_requests.get(request.request_id) + if (prev is not None and prev.remote_matched_token_num + and prev.block_ids_per_group): + # The request already went through an external load (blocks were + # allocated) and returned to WAITING -- a KV load failure or a + # preemption. The manager may still advertise blocks whose storage + # is gone, so re-matching risks an endless fail-reschedule loop; + # recompute locally instead. (A pending re-query after a failed + # allocation has empty block_ids_per_group and is not affected.) + logger.warning("req:%s re-queried after an external load attempt, " + "skip external match", request.request_id) + prev.local_matched_token_num = num_computed_tokens + prev.remote_matched_token_num = 0 + prev.has_saved_block_num = num_computed_tokens // self._manager_block_size + return 0, False + computed_blocks = num_computed_tokens // self._manager_block_size is_query_done, need_load_locations = ( From 924a9521cafbf8eacec0d4c099cbca656acda0ce Mon Sep 17 00:00:00 2001 From: xiaozeyu Date: Tue, 28 Jul 2026 18:18:19 +0800 Subject: [PATCH 6/7] [integration_test] harden vllm e2e harness and fix hybrid prompt coverage Harness hardening (task B) plus the P2 coverage fix: * Hybrid prompts were ~578 tokens = 1 x 528 manager block, so multi-block state mapping and incremental save never ran under hybrid. make_base_prompts now emits 140 sentences (~2100 tokens > 3 x 528) for hybrid models. * wait_for_captures timeout raises AssertionError instead of warning. * Expected ref/loaded capture counts are computed per prompt from the actual tokenization (len // mbs, shared-prefix for loads) and asserted as exact lower bounds via assert_report_ok(min_matched=...). * bit-exact comparison is the default; cosine fallback only with KVCM_E2E_ALLOW_COSINE=1. * compare_captures iterates the loaded capture's layers (a loaded layer without a reference is a hard error); mamba align-mode null-state layers may legitimately be absent on either side and test_connector skips capturing null (id 0) state blocks, mirroring the connector's vacuous transfer. * VllmServer/ScenarioEnv: connector_name selection (mutation meta-test), log_level, kv_transfer extra-config overrides, kv_load_failure_policy, key_count_per_file, and a ScenarioEnv helper owning manager/vLLM lifecycle for the new custom scenarios; send_completions accepts token-id prompts and extra payload fields; file-backend root_path gets a trailing slash so block files land inside the dir. * test_connector: drop a duplicated capture line; add MutatedConnector (slot -1 off-by-one, test-side injection only) for the mutation meta-test. --- integration_test/vllm_e2e/e2e_lib.py | 367 +++++++++++++++----- integration_test/vllm_e2e/test_connector.py | 32 +- 2 files changed, 312 insertions(+), 87 deletions(-) diff --git a/integration_test/vllm_e2e/e2e_lib.py b/integration_test/vllm_e2e/e2e_lib.py index ae7698862..29190eda8 100644 --- a/integration_test/vllm_e2e/e2e_lib.py +++ b/integration_test/vllm_e2e/e2e_lib.py @@ -43,9 +43,17 @@ import requests logger = logging.getLogger("vllm_e2e") +# This module only runs inside test drivers; make the orchestration evidence +# (block counts, verification report, log-scan results) visible in test.log. +logging.basicConfig( + level=logging.INFO, + format="%(asctime)s %(name)s %(levelname)s %(message)s") MODEL_PATH = os.environ.get("KVCM_E2E_MODEL", "/root/ws/resources/models/Qwen2.5-7B-Instruct") COSINE_THRESHOLD = 0.9999 +# Bit-exact comparison is the default (empirically all scenarios achieve it). +# Cosine fallback must be explicitly requested. +ALLOW_COSINE = os.environ.get("KVCM_E2E_ALLOW_COSINE", "0") == "1" def is_hybrid_model(model_path: str) -> bool: @@ -137,7 +145,7 @@ def wait_http(url: str, timeout: float, post_body: Optional[dict] = None) -> boo # KVCM manager # --------------------------------------------------------------------------- # class ManagerProcess: - def __init__(self, workdir: str, storage_root: str): + def __init__(self, workdir: str, storage_root: str, key_count_per_file: int = 8): self.workdir = workdir os.makedirs(workdir, exist_ok=True) self.rpc_port = free_port() @@ -145,6 +153,7 @@ def __init__(self, workdir: str, storage_root: str): self.admin_rpc_port = free_port() self.admin_http_port = free_port() self.storage_root = storage_root + self.key_count_per_file = key_count_per_file self.proc: Optional[subprocess.Popen] = None self.config_path = os.path.join(workdir, "startup_config.json") @@ -157,8 +166,10 @@ def _write_config(self): "type": "file", "global_unique_name": "nfs_01", "storage_spec": { - "root_path": self.storage_root, - "key_count_per_file": 8, + # The backend concatenates root_path + key with no + # separator; the trailing slash keeps files inside the dir. + "root_path": self.storage_root.rstrip("/") + "/", + "key_count_per_file": self.key_count_per_file, }, }, "instance_group": { @@ -247,7 +258,11 @@ class VllmServer: def __init__(self, workdir: str, capture_dir: str, manager_uri: str, tp_size: int, coordinator_base_port: int, instance_id: str, preferred_block_size: int, - enable_prefix_caching: bool): + enable_prefix_caching: bool, + connector_name: str = "VerifyingConnector", + log_level: str = "INFO", + extra_config_overrides: Optional[dict] = None, + kv_load_failure_policy: Optional[str] = None): self.workdir = workdir os.makedirs(workdir, exist_ok=True) self.capture_dir = capture_dir @@ -259,6 +274,10 @@ def __init__(self, workdir: str, capture_dir: str, manager_uri: str, self.instance_id = instance_id self.preferred_block_size = preferred_block_size self.enable_prefix_caching = enable_prefix_caching + self.connector_name = connector_name + self.log_level = log_level + self.extra_config_overrides = extra_config_overrides or {} + self.kv_load_failure_policy = kv_load_failure_policy self.proc: Optional[subprocess.Popen] = None def base_url(self) -> str: @@ -271,14 +290,17 @@ def start(self, repo_root: str, log_suffix: str = ""): "instance_group": "default", "instance_id": self.instance_id, "preferred_block_size": self.preferred_block_size, - "log_level": "INFO", + "log_level": self.log_level, } + extra_config.update(self.extra_config_overrides) kv_transfer_config = { - "kv_connector": "VerifyingConnector", + "kv_connector": self.connector_name, "kv_role": "kv_both", "kv_connector_module_path": "test_connector", "kv_connector_extra_config": extra_config, } + if self.kv_load_failure_policy: + kv_transfer_config["kv_load_failure_policy"] = self.kv_load_failure_policy cmd = [ find_python(), "-m", "vllm.entrypoints.openai.api_server", "--model", MODEL_PATH, @@ -345,13 +367,43 @@ def tokenize(base_url: str, prompt: str) -> list[int]: return r.json()["tokens"] +def get_manager_block_size(manager_uri: str, instance_id: str) -> int: + """Ask the manager for the registered instance's manager block size.""" + r = requests.post(f"{manager_uri}/api/getInstanceInfo", + json={"trace_id": "e2e_bs", "instance_id": instance_id}, + timeout=10) + r.raise_for_status() + block_size = r.json()["instance_info"]["block_size"] + assert block_size > 0, f"bad manager block size: {block_size}" + return block_size + + +def block_token_hash(token_ids: list[int]) -> str: + """Token-content hash used in capture file names (mirrors test_connector).""" + import hashlib + import torch + return hashlib.sha256( + torch.tensor(token_ids, dtype=torch.int64).numpy().tobytes() + ).hexdigest()[:16] + + +def full_block_hashes(token_ids: list[int], manager_block_size: int) -> list[str]: + """Per-manager-block capture hashes for the full blocks of a token stream.""" + n = len(token_ids) // manager_block_size + return [ + block_token_hash(token_ids[i * manager_block_size:(i + 1) * manager_block_size]) + for i in range(n) + ] + + def wait_for_prefix_cached(manager_uri: str, instance_id: str, - token_ids: list[int], timeout: float = 120.0) -> bool: - """Poll the manager until the prefix for token_ids is cached (committed). + token_ids: list[int], min_blocks: int, + timeout: float = 120.0) -> bool: + """Poll the manager until at least min_blocks of the prefix are committed. Mirrors what the connector's get_num_new_matched_tokens queries (query_type=QT_PREFIX_MATCH, block_mask offset=0 for a fresh request), so it - guarantees phase 2 will actually hit the external cache. + guarantees phase 2 will actually hit the external cache for all min_blocks. """ deadline = time.time() + timeout payload = { @@ -369,7 +421,7 @@ def wait_for_prefix_cached(manager_uri: str, instance_id: str, data = r.json() if data.get("header", {}).get("status", {}).get("code") == "OK": locs = data.get("locations", []) - if locs: + if len(locs) >= min_blocks: logger.info("prefix cached: %d location(s)", len(locs)) return True except Exception: @@ -379,19 +431,24 @@ def wait_for_prefix_cached(manager_uri: str, instance_id: str, return False -def send_completions(base_url: str, prompts: list[str], max_tokens: int = 4, - temperature: float = 0.0) -> list[dict]: - """Send prompts concurrently and return the OpenAI responses.""" +def send_completions(base_url: str, prompts: list, max_tokens: int = 4, + temperature: float = 0.0, **extra_payload) -> list[dict]: + """Send prompts concurrently and return the OpenAI responses. + + Each prompt may be a string or a list of token ids (the completions API + accepts both). extra_payload is merged into the request body (e.g. + return_token_ids=True).""" from concurrent.futures import ThreadPoolExecutor client_url = f"{base_url}/v1/completions" - def _one(prompt: str) -> dict: + def _one(prompt) -> dict: payload = { "model": "qwen", "prompt": prompt, "max_tokens": max_tokens, "temperature": temperature, + **extra_payload, } r = requests.post(client_url, json=payload, timeout=300) r.raise_for_status() @@ -409,7 +466,13 @@ def count_captures(capture_dir: str, kind: str) -> int: def wait_for_captures(capture_dir: str, kind: str, expected: int, - timeout: float = 120.0): + timeout: float = 120.0) -> int: + """Wait until at least ``expected`` captures of ``kind`` exist. + + Raises AssertionError on timeout: a missing capture means the connector + never exercised the code path under test, so the scenario must fail rather + than silently verify fewer blocks. + """ deadline = time.time() + timeout while time.time() < deadline: n = count_captures(capture_dir, kind) @@ -418,9 +481,8 @@ def wait_for_captures(capture_dir: str, kind: str, expected: int, return n time.sleep(1.0) n = count_captures(capture_dir, kind) - logger.warning("timed out waiting for %s captures: got %d, want %d", - kind, n, expected) - return n + raise AssertionError( + f"timed out waiting for {kind} captures: got {n}, want {expected}") def _cosine(a, b) -> float: @@ -461,6 +523,7 @@ def compare_captures(capture_dir: str, tp_size: int) -> dict: "cosine_pass": 0, "failures": [], "loaded_without_ref": [], + "matched_keys": [], } for key, loaded_path in sorted(loaded.items()): @@ -474,8 +537,16 @@ def compare_captures(capture_dir: str, tp_size: int) -> dict: all_bit_exact = True worst_cosine = 1.0 - for layer_name, ref_kv in ref["kv"].items(): - got_kv = got["kv"][layer_name] + # Compare every layer present in the *loaded* capture: each one was + # actually written by the connector and must match its reference. + # Mamba "align" state layers can legitimately be absent on either side + # (vLLM materializes states only at segment boundaries; interior blocks + # get the null block and the connector transfers them vacuously) -- but + # a loaded layer without a reference is a hard error. + for layer_name, got_kv in got["kv"].items(): + assert layer_name in ref["kv"], ( + f"loaded layer {layer_name} of {key} has no reference capture") + ref_kv = ref["kv"][layer_name] # Attention groups are a single Tensor; mamba/linear/gdn groups are a # list[Tensor] (e.g. [conv_state, ssm_state]). Compare uniformly. if isinstance(ref_kv, (list, tuple)): @@ -497,7 +568,7 @@ def compare_captures(capture_dir: str, tp_size: int) -> dict: all_bit_exact = False cos = _cosine(ref_t, got_t) worst_cosine = min(worst_cosine, cos) - if cos < COSINE_THRESHOLD: + if not ALLOW_COSINE or cos < COSINE_THRESHOLD: report["failures"].append({ "key": key, "layer": f"{layer_name}[{si}]", @@ -505,6 +576,7 @@ def compare_captures(capture_dir: str, tp_size: int) -> dict: }) report["matched"] += 1 + report["matched_keys"].append(key) if all_bit_exact: report["bit_exact"] += 1 else: @@ -515,16 +587,25 @@ def compare_captures(capture_dir: str, tp_size: int) -> dict: return report -def assert_report_ok(report: dict): +def assert_report_ok(report: dict, min_matched: int = 1): + """Assert the comparison succeeded. + + min_matched is the exact lower bound of (ref, loaded) capture pairs computed + from the prompts' tokenization (num full manager blocks x tp ranks); a lower + count means some blocks were silently never saved or never loaded. + """ problems = [] if report["loaded_without_ref"]: problems.append( f"loaded captures with no matching reference: {report['loaded_without_ref']}" ) if report["failures"]: - problems.append(f"cosine failures: {report['failures']}") - if report["matched"] == 0: - problems.append("no loaded captures were matched against references") + kind = "cosine" if ALLOW_COSINE else "bit-exact" + problems.append(f"{kind} failures: {report['failures']}") + if report["matched"] < min_matched: + problems.append( + f"matched {report['matched']} loaded captures, expected >= {min_matched}" + ) if problems: raise AssertionError("KV verification failed: " + "; ".join(problems)) logger.info( @@ -537,100 +618,216 @@ def assert_report_ok(report: dict): # --------------------------------------------------------------------------- # # Scenario runner # --------------------------------------------------------------------------- # +class ScenarioEnv: + """Owns one scenario's manager + vLLM server lifecycle and scratch dirs. + + Custom scenarios (partial-hit, full-hit, load-failure, multi-turn) share + this; run_e2e keeps its own two-phase flow on top of the same pieces. + """ + + def __init__(self, scenario: str, tp_size: int = 1, + preferred_block_size: int = 0, + enable_prefix_caching: Optional[bool] = None, + connector_name: str = "VerifyingConnector", + log_level: str = "INFO", + extra_config_overrides: Optional[dict] = None, + key_count_per_file: int = 8, + kv_load_failure_policy: Optional[str] = None): + self.scenario = scenario + self.tp_size = tp_size + self.hybrid = is_hybrid_model(MODEL_PATH) + self.preferred_block_size = 0 if self.hybrid else preferred_block_size + self.enable_prefix_caching = (self.hybrid if enable_prefix_caching is None + else enable_prefix_caching) + self.connector_name = connector_name + self.log_level = log_level + self.extra_config_overrides = extra_config_overrides + self.kv_load_failure_policy = kv_load_failure_policy + + self.repo_root = find_repo_root() + scratch_root = (os.environ.get("TEST_TMPDIR") + or os.environ.get("TMPDIR") or "/tmp") + self.base_workdir = os.path.join(scratch_root, "kvcm_vllm_e2e", scenario) + if os.path.exists(self.base_workdir): + shutil.rmtree(self.base_workdir) + self.storage_root = os.path.join(self.base_workdir, "nfs") + self.capture_dir = os.path.join(self.base_workdir, "captures") + self.vllm_dir = os.path.join(self.base_workdir, "vllm") + os.makedirs(self.storage_root, exist_ok=True) + + self.instance_id = f"e2e-{scenario}-{uuid.uuid4().hex[:8]}" + self.manager = ManagerProcess( + os.path.join(self.base_workdir, "manager"), self.storage_root, + key_count_per_file=key_count_per_file) + self.vllm: Optional[VllmServer] = None + + def start_manager(self): + self.manager.start(self.repo_root) + + def start_vllm(self, log_suffix: str = "") -> VllmServer: + self.vllm = VllmServer( + self.vllm_dir, self.capture_dir, self.manager.manager_uri(), + self.tp_size, coordinator_base_port=free_port(), + instance_id=self.instance_id, + preferred_block_size=self.preferred_block_size, + enable_prefix_caching=self.enable_prefix_caching, + connector_name=self.connector_name, + log_level=self.log_level, + extra_config_overrides=self.extra_config_overrides, + kv_load_failure_policy=self.kv_load_failure_policy, + ) + self.vllm.start(self.repo_root, log_suffix=log_suffix) + return self.vllm + + def restart_vllm(self, log_suffix: str = "") -> VllmServer: + """Restart vLLM to drop its local prefix cache (KVCM state persists).""" + if self.vllm: + self.vllm.stop() + return self.start_vllm(log_suffix) + + def manager_block_size(self) -> int: + return get_manager_block_size(self.manager.manager_uri(), self.instance_id) + + def scan_connector_logs(self, pattern: str) -> list: + """Regex-scan all vLLM std streams; returns list of match groups.""" + import re + out = [] + for path in glob.glob(os.path.join(self.vllm_dir, "vllm*.std*")): + with open(path, errors="replace") as f: + for line in f: + m = re.search(pattern, line) + if m: + out.append(m.groups() if m.groups() else m.group(0)) + return out + + def stop(self): + if self.vllm: + self.vllm.stop() + self.manager.stop() + + +def make_base_prompts(num_prompts: int, hybrid: bool) -> list[str]: + """Distinct, deterministic prompts. Each sentence carries a unique counter + so every manager block has unique token content -- this avoids hash + collisions between blocks with identical text but different KV (RoPE is + position-dependent). + + Hybrid models pin the manager block size to the scheduler block size (528), + so their prompts must be much longer to span multiple manager blocks + (empirically 140 sentences ~ 2100 tokens > 3 x 528). Full-attention models + use manager blocks of 16/32 tokens, where 40 sentences (~580 tokens) already + span dozens of blocks. + """ + num_sentences = 140 if hybrid else 40 + return [ + f"Prompt number {i}. " + " ".join( + f"Sentence {j} of prompt {i} has value {j * 7 + i * 131}." + for j in range(num_sentences) + ) + for i in range(num_prompts) + ] + + +def shared_token_prefix_len(a: list[int], b: list[int]) -> int: + n = 0 + for x, y in zip(a, b): + if x != y: + break + n += 1 + return n + + def run_e2e(scenario: str, tp_size: int, num_prompts: int, - preferred_block_size: int): + preferred_block_size: int, connector_name: str = "VerifyingConnector", + expect_verification_failure: bool = False): """Run one full save-then-load verification scenario. Full-attention models: prefix caching off, one server across both phases. Hybrid models: prefix caching on (align mode -> per-group block tables), the vLLM server is restarted between phases so phase 2 loads from KVCM instead of hitting the local prefix cache. + + connector_name selects the connector class inside test_connector.py; the + mutation meta-test passes "MutatedConnector" and sets + expect_verification_failure=True to prove the harness detects an injected + off-by-one in the token translation. """ import torch # noqa: F401 (ensure torch importable early for clear errors) - hybrid = is_hybrid_model(MODEL_PATH) - # Hybrid mamba state is per scheduler block, so preferred_block_size can only - # differ from the vLLM block size for pure-attention models. - if hybrid: - preferred_block_size = 0 - - repo_root = find_repo_root() - scratch_root = os.environ.get("TEST_TMPDIR") or os.environ.get("TMPDIR") or "/tmp" - base_workdir = os.path.join(scratch_root, "kvcm_vllm_e2e", scenario) - if os.path.exists(base_workdir): - shutil.rmtree(base_workdir) - storage_root = os.path.join(base_workdir, "nfs") - manager_dir = os.path.join(base_workdir, "manager") - vllm_dir = os.path.join(base_workdir, "vllm") - capture_dir = os.path.join(base_workdir, "captures") - os.makedirs(storage_root, exist_ok=True) - - instance_id = f"e2e-{scenario}-{uuid.uuid4().hex[:8]}" + env = ScenarioEnv(scenario, tp_size=tp_size, + preferred_block_size=preferred_block_size, + connector_name=connector_name) + hybrid = env.hybrid + capture_dir = env.capture_dir logger.info("scenario=%s model=%s hybrid=%s tp=%d prompts=%d preferred_bs=%d", - scenario, MODEL_PATH, hybrid, tp_size, num_prompts, preferred_block_size) - - manager = ManagerProcess(manager_dir, storage_root) + scenario, MODEL_PATH, hybrid, tp_size, num_prompts, + env.preferred_block_size) - def make_server(coordinator_port): - return VllmServer( - vllm_dir, capture_dir, manager.manager_uri(), tp_size, - coordinator_base_port=coordinator_port, - instance_id=instance_id, - preferred_block_size=preferred_block_size, - enable_prefix_caching=hybrid, - ) + try: + env.start_manager() + vllm = env.start_vllm(log_suffix="" if not hybrid else "_p1") - vllm = make_server(free_port()) + base_prompts = make_base_prompts(num_prompts, hybrid) + suffixes = [f" Now answer question {i}: what is 2+2?" for i in range(num_prompts)] + phase2_prompts = [p + s for p, s in zip(base_prompts, suffixes)] - try: - manager.start(repo_root) - vllm.start(repo_root, log_suffix="" if not hybrid else "_p1") - - # Distinct, deterministic prompts. Each sentence carries a unique counter - # so every manager block has unique token content -- this avoids hash - # collisions between blocks with identical text but different KV (RoPE is - # position-dependent). Long enough to span several manager blocks so the - # translation layer is actually exercised. - base_prompts = [ - f"Prompt number {i}. " + " ".join( - f"Sentence {j} of prompt {i} has value {j * 7 + i * 131}." - for j in range(40) - ) - for i in range(num_prompts) + # Compute per-prompt expected block counts from the actual tokenization + # so a silently dropped prompt (or block) fails the run. + mbs = env.manager_block_size() + base_tokens = [tokenize(vllm.base_url(), p) for p in base_prompts] + phase2_tokens = [tokenize(vllm.base_url(), p) for p in phase2_prompts] + expected_save_blocks = [len(t) // mbs for t in base_tokens] + # Loads cover the shared token prefix (the base/suffix boundary token may + # re-merge under tokenization, shortening the shared prefix by one). + expected_load_blocks = [ + min(shared_token_prefix_len(b, p2) // mbs, s) + for b, p2, s in zip(base_tokens, phase2_tokens, expected_save_blocks) ] - suffixes = [f" Now answer question {i}: what is 2+2?" for i in range(num_prompts)] + assert all(n >= 1 for n in expected_load_blocks), ( + f"prompts too short to span a manager block (mbs={mbs}): " + f"{expected_load_blocks}") + logger.info("mbs=%d expected save blocks=%s load blocks=%s", + mbs, expected_save_blocks, expected_load_blocks) # ---- Phase 1: fresh prefill -> connector saves -> reference capture. logger.info("phase 1: sending %d fresh prompts", num_prompts) send_completions(vllm.base_url(), base_prompts) - wait_for_captures(capture_dir, "ref", expected=num_prompts, timeout=180) + wait_for_captures(capture_dir, "ref", + expected=tp_size * sum(expected_save_blocks), timeout=180) # The save is committed to the manager asynchronously after the ref # capture (which fires when the save is submitted). Wait until the manager # actually has the prefix, otherwise phase 2 would find no match. - phase2_prompts = [p + s for p, s in zip(base_prompts, suffixes)] - for p in base_prompts: - toks = tokenize(vllm.base_url(), p) - if not wait_for_prefix_cached(manager.manager_uri(), instance_id, toks): + for toks, blocks in zip(base_tokens, expected_save_blocks): + if not wait_for_prefix_cached(env.manager.manager_uri(), env.instance_id, + toks, min_blocks=blocks): raise AssertionError("save was not committed to the manager in time") # Hybrid models keep prefix caching on, which also populates the local # prefix cache; restart vLLM so phase 2 loads from KVCM, not locally. if hybrid: logger.info("restarting vLLM before phase 2 (clear local prefix cache)") - vllm.stop() - vllm = make_server(free_port()) - vllm.start(repo_root, log_suffix="_p2") + vllm = env.restart_vllm(log_suffix="_p2") # ---- Phase 2: same prefix + suffix -> connector loads -> loaded capture. logger.info("phase 2: sending %d prefix+suffix prompts", num_prompts) send_completions(vllm.base_url(), phase2_prompts) - wait_for_captures(capture_dir, "loaded", expected=num_prompts, timeout=180) + wait_for_captures(capture_dir, "loaded", + expected=tp_size * sum(expected_load_blocks), timeout=180) report = compare_captures(capture_dir, tp_size) - assert_report_ok(report) + min_matched = tp_size * sum(expected_load_blocks) + if expect_verification_failure: + try: + assert_report_ok(report, min_matched=min_matched) + except AssertionError as e: + logger.info("verification failed as expected: %s", e) + return + raise AssertionError( + "mutated connector passed KV verification; the harness is blind") + assert_report_ok(report, min_matched=min_matched) logger.info("scenario %s PASSED: %s", scenario, json.dumps( - {k: v for k, v in report.items() if k != "failures"}, default=str)) + {k: v for k, v in report.items() if k not in ("failures", "matched_keys")}, + default=str)) finally: - vllm.stop() - manager.stop() + env.stop() diff --git a/integration_test/vllm_e2e/test_connector.py b/integration_test/vllm_e2e/test_connector.py index 42d71c0df..4d5055306 100644 --- a/integration_test/vllm_e2e/test_connector.py +++ b/integration_test/vllm_e2e/test_connector.py @@ -267,12 +267,16 @@ def _capture_block(self, kind, token_ids, block_ids_per_group, manager_block_idx flat = kv_cache.permute(0, 2, 1, 3).reshape(-1, per_token) gathered = flat[slot_tensor, :].contiguous() # [n_tok, per_token] kv_by_layer[layer_name] = gathered.cpu() - kv_by_layer[layer_name] = gathered.cpu() else: # State stored once per group block; the manager block's last - # token selects the block (mirrors _state_block_ids). + # token selects the block (mirrors _state_block_ids). vLLM's + # mamba "align" mode materializes states only at segment + # boundaries -- interior blocks hold the null block (id 0) and + # carry no state to capture (the connector skips them too). logical = ((manager_block_idx + 1) * mbs - 1) // group_bs block_id = block_table[logical] + if block_id == 0: + continue for layer_name in layer_names: states = self._kv_caches[layer_name] # list[Tensor] kv_by_layer[layer_name] = [s[block_id].detach().cpu() for s in states] @@ -285,3 +289,27 @@ def _capture_block(self, kind, token_ids, block_ids_per_group, manager_block_idx logger.warning( "VerifyingConnector captured %s block=%d tokens=%d..%d tp=%s -> %s", kind, manager_block_idx, positions[0], positions[-1], self._tp_rank, path) + + +class MutatedConnector(VerifyingConnector): + """Meta-test connector: injects an off-by-one into the attention token + translation (every gathered/scattered slot shifted by -1). + + The shift is symmetric between save and load, so with contiguous block + tables a transport round trip cancels it in the interior of the loaded + range (slot(t)-1 == slot(t-1)); the leak is at the boundary: the last + loaded token's true slot is never written and keeps stale (uninitialized) + data. The capture-based verification reads the cache through vLLM's own + slot mapping and must observe that divergence -- the mutation e2e test + asserts that verification FAILS with this connector. + + -1 (not +1) keeps every shifted slot in bounds: vLLM reserves physical + block 0 as the null block, so real slots are >= kernel_block_size and + slot-1 >= 0, while slot+1 of the cache's last block would read/write out + of bounds. Only reachable through the test-side ``kv_connector_module_path`` + injection; never part of the production wheel. + """ + + def _attn_token_indices(self, group, manager_block_idxes, block_table): + out = super()._attn_token_indices(group, manager_block_idxes, block_table) + return [[slot - 1 for slot in block] for block in out] From 71fa2d5635b4a6b95425aa366a436cda470693e9 Mon Sep 17 00:00:00 2001 From: xiaozeyu Date: Tue, 28 Jul 2026 18:18:33 +0800 Subject: [PATCH 7/7] [integration_test] add mutation meta-test and four vllm e2e scenarios * test_mutation (B4): runs the basic scenario with MutatedConnector (slot -1 in _attn_token_indices) and asserts KV verification FAILS -- proof the capture-based harness catches symmetric translation bugs and is not vacuous. * test_full_hit (C2): prompt trimmed to an exact multiple of the manager block size, resent after being fully saved. Regression for the synchronous full-hit crash (vllm 0.26.0 scheduler.py 'assert num_new_tokens > 0'); asserts the engine survives and 0 < matched < prompt tokens. * test_partial_hit (C1): staged A / A+B / A+B+C prompts with prefix caching on; exercises the non-zero-offset incremental manager query and the incremental save extension, asserts a logged query with offset > 0 and verifies the blocks saved through the incremental path. * test_load_failure (C3): deletes the tail half of the per-block storage files between save and load (key_count_per_file=1, block_per_load_task=1, kv_load_failure_policy=recompute). Full-attn: failures reported to vLLM, surviving head blocks verify bit-exact, mismatches confined to deleted blocks. Hybrid: failure swallowed by design (vLLM invalid-block recovery is single-group only), asserts no hang/crash and the failure log. * test_multi_turn (C4): turn 1 decodes past a manager block boundary (ignore_eos + return_token_ids), asserts the manager committed more blocks than the prompt covers; turn 2 embeds turn 1 prompt+output as token ids and must externally match beyond prompt-only coverage with verified KV. All scenarios pass for both Qwen2.5-7B-Instruct (full-attn) and Qwen3.5-4B (hybrid) alongside the original basic/concurrent/tp regressions. --- integration_test/vllm_e2e/BUILD | 82 +++++++++ integration_test/vllm_e2e/test_full_hit.py | 82 +++++++++ .../vllm_e2e/test_load_failure.py | 161 ++++++++++++++++++ integration_test/vllm_e2e/test_multi_turn.py | 112 ++++++++++++ integration_test/vllm_e2e/test_mutation.py | 35 ++++ integration_test/vllm_e2e/test_partial_hit.py | 96 +++++++++++ 6 files changed, 568 insertions(+) create mode 100644 integration_test/vllm_e2e/test_full_hit.py create mode 100644 integration_test/vllm_e2e/test_load_failure.py create mode 100644 integration_test/vllm_e2e/test_multi_turn.py create mode 100644 integration_test/vllm_e2e/test_mutation.py create mode 100644 integration_test/vllm_e2e/test_partial_hit.py diff --git a/integration_test/vllm_e2e/BUILD b/integration_test/vllm_e2e/BUILD index c025478cd..f8ceef86b 100644 --- a/integration_test/vllm_e2e/BUILD +++ b/integration_test/vllm_e2e/BUILD @@ -59,3 +59,85 @@ py_test( timeout = "eternal", deps = [":e2e_lib"], ) + +py_test( + name = "test_full_hit", + srcs = ["test_full_hit.py"], + data = [ + "//kv_cache_manager:kv_cache_manager_bin", + ], + imports = ["."], + tags = [ + "no-remote-exec", + "gpu", # requires 1+ GPU + "exclusive", # GPU tests must run serially to avoid CUDA OOM contention + ], + timeout = "eternal", + deps = [":e2e_lib"], +) + +py_test( + name = "test_partial_hit", + srcs = ["test_partial_hit.py"], + data = [ + "//kv_cache_manager:kv_cache_manager_bin", + ], + imports = ["."], + tags = [ + "no-remote-exec", + "gpu", # requires 1+ GPU + "exclusive", # GPU tests must run serially to avoid CUDA OOM contention + ], + timeout = "eternal", + deps = [":e2e_lib"], +) + +py_test( + name = "test_load_failure", + srcs = ["test_load_failure.py"], + data = [ + "//kv_cache_manager:kv_cache_manager_bin", + ], + imports = ["."], + tags = [ + "no-remote-exec", + "gpu", # requires 1+ GPU + "exclusive", # GPU tests must run serially to avoid CUDA OOM contention + ], + timeout = "eternal", + deps = [":e2e_lib"], +) + +py_test( + name = "test_multi_turn", + srcs = ["test_multi_turn.py"], + data = [ + "//kv_cache_manager:kv_cache_manager_bin", + ], + imports = ["."], + tags = [ + "no-remote-exec", + "gpu", # requires 1+ GPU + "exclusive", # GPU tests must run serially to avoid CUDA OOM contention + ], + timeout = "eternal", + deps = [":e2e_lib"], +) + +# Meta-test: injects an off-by-one into the connector's token translation and +# asserts the KV verification FAILS -- proof the harness is not vacuous. +py_test( + name = "test_mutation", + srcs = ["test_mutation.py"], + data = [ + "//kv_cache_manager:kv_cache_manager_bin", + ], + imports = ["."], + tags = [ + "no-remote-exec", + "gpu", # requires 1+ GPU + "exclusive", # GPU tests must run serially to avoid CUDA OOM contention + ], + timeout = "eternal", + deps = [":e2e_lib"], +) diff --git a/integration_test/vllm_e2e/test_full_hit.py b/integration_test/vllm_e2e/test_full_hit.py new file mode 100644 index 000000000..39c214ecb --- /dev/null +++ b/integration_test/vllm_e2e/test_full_hit.py @@ -0,0 +1,82 @@ +"""test_full_hit: full-prompt external hit must not crash the engine. + +Regression test for the synchronous-load full-hit bug: this connector reports +external matches with load_kv_async=False, so vLLM schedules +``num_tokens - num_computed_tokens`` new tokens and asserts that count is > 0 +(vllm/v1/core/sched/scheduler.py, waiting-queue loop: ``assert num_new_tokens +> 0``). Without capping, a prompt whose token count is an exact multiple of the +manager block size and whose blocks are all externally cached would make the +count 0 and kill the engine. + +Phase 1 saves a prompt of exactly N manager blocks; phase 2 resends the very +same prompt (as explicit token ids, so tokenization cannot shift the length). +Asserts: the engine survives, the completion is well-formed, and the connector +reports 0 < matched < prompt tokens (the cap dropped at least the last block). + +Runs against both full-attention and hybrid models via $KVCM_E2E_MODEL. +""" + +import logging +import unittest + +from e2e_lib import ( + ScenarioEnv, make_base_prompts, send_completions, tokenize, + wait_for_prefix_cached, +) + +logger = logging.getLogger("vllm_e2e") + + +class TestFullHit(unittest.TestCase): + def test_full_hit(self): + env = ScenarioEnv("full_hit") + try: + env.start_manager() + vllm = env.start_vllm(log_suffix="_p1" if env.hybrid else "") + + mbs = env.manager_block_size() + # Trim a long-enough prompt's token ids to an exact multiple of the + # manager block size (>= 2 blocks so the cap has room to drop one). + toks = tokenize(vllm.base_url(), make_base_prompts(1, env.hybrid)[0]) + num_blocks = len(toks) // mbs + self.assertGreaterEqual( + num_blocks, 2, f"prompt too short: {len(toks)} tokens, mbs={mbs}") + prompt_ids = toks[:num_blocks * mbs] + logger.info("full-hit prompt: %d tokens = %d x %d", + len(prompt_ids), num_blocks, mbs) + + # Phase 1: fresh prefill -> all blocks saved. + resp1 = send_completions(vllm.base_url(), [prompt_ids])[0] + self.assertTrue(resp1["choices"][0]["text"]) + self.assertTrue(wait_for_prefix_cached( + env.manager.manager_uri(), env.instance_id, prompt_ids, + min_blocks=num_blocks)) + + if env.hybrid: + # Clear the local prefix cache so phase 2 goes external. + vllm = env.restart_vllm(log_suffix="_p2") + + # Phase 2: the exact same prompt -> full external hit. Without the + # cap this crashes the engine (assert num_new_tokens > 0). + resp2 = send_completions(vllm.base_url(), [prompt_ids])[0] + self.assertTrue(resp2["choices"][0]["text"]) + + # The engine must still be alive and serving. + resp3 = send_completions(vllm.base_url(), ["sanity check prompt"])[0] + self.assertTrue(resp3["choices"][0]["text"]) + + # Connector-side evidence: matched > 0 (external hit happened) and + # matched < prompt tokens (the cap left tokens to recompute). + matched = [int(g[0]) for g in + env.scan_connector_logs(r"matched (\d+) external tokens")] + self.assertTrue(matched, "no 'matched N external tokens' log found") + hit = [m for m in matched if m > 0] + self.assertTrue(hit, f"no positive external match in {matched}") + self.assertTrue(all(m < len(prompt_ids) for m in hit), + f"match not capped below prompt len: {matched}") + finally: + env.stop() + + +if __name__ == "__main__": + unittest.main() diff --git a/integration_test/vllm_e2e/test_load_failure.py b/integration_test/vllm_e2e/test_load_failure.py new file mode 100644 index 000000000..e353dbfe7 --- /dev/null +++ b/integration_test/vllm_e2e/test_load_failure.py @@ -0,0 +1,161 @@ +"""test_load_failure: storage loss between save and load must degrade, not kill. + +Phase 1 saves a long prompt; the test then deletes the storage files of the +*tail half* of the manager blocks (resolved through the manager's ordered +getCacheLocation response, one file per block via key_count_per_file=1). +Phase 2 reloads the same prefix with ``block_per_load_task=1`` so every block +fails or succeeds independently. + +Full-attention models (single group, report_failures=True): the connector +reports the failed blocks' vLLM block ids; with +``kv_load_failure_policy="recompute"`` (vLLM 0.26.0 defaults to "fail", which +turns any load failure into a 500) vLLM truncates the computed-token count at +the first invalid block and recomputes from there +(vllm/v1/core/sched/scheduler.py::_handle_invalid_blocks / +_update_requests_with_invalid_blocks). Asserts: the request returns a normal +completion, the failures were logged and reported, every *surviving* head +block's loaded KV is bit-exact, and any mismatching capture belongs to a +deleted (recomputed) block. Recomputed blocks are not held to bit-exactness: +they contain freshly recomputed KV whose numerics depend on prefill chunking, +which is vLLM's business, not the connector's. + +Hybrid models (multiple groups, report_failures=False): vLLM's invalid-block +recovery only supports single-group block tables, so the connector only logs +the failure. Asserts: the request still returns (no hang, no crash) and the +failure was logged. KV content is NOT verified: with the failure swallowed +the affected blocks keep garbage by design. + +This scenario also regression-tests the fail-reschedule loop fix: a request +whose external load failed must not re-match external blocks on requeue +(v1_connector.get_num_new_matched_tokens retry guard), otherwise the engine +loops load-fail-reschedule forever and the request hangs. +""" + +import logging +import os +import unittest +from urllib.parse import urlparse + +import requests + +from e2e_lib import ( + ScenarioEnv, compare_captures, full_block_hashes, make_base_prompts, + send_completions, tokenize, wait_for_captures, wait_for_prefix_cached, +) + +logger = logging.getLogger("vllm_e2e") + + +def get_block_files(manager_uri: str, instance_id: str, token_ids: list[int], + spec_name: str = "tp0_g0") -> list[str]: + """Per-manager-block storage file paths, in block order, from the manager's + getCacheLocation response.""" + r = requests.post(f"{manager_uri}/api/getCacheLocation", json={ + "trace_id": "e2e_block_files", + "token_ids": token_ids, + "instance_id": instance_id, + "query_type": "QT_PREFIX_MATCH", + "block_mask": {"offset": 0}, + }, timeout=10) + r.raise_for_status() + files = [] + for location in r.json().get("locations", []): + for spec in location.get("location_specs", []): + if spec["name"] == spec_name: + # uri: file://?size=... + files.append(urlparse(spec["uri"]).path) + return files + + +class TestLoadFailure(unittest.TestCase): + def test_load_failure(self): + env = ScenarioEnv( + "load_failure", + extra_config_overrides={"block_per_load_task": 1}, + key_count_per_file=1, # one file per block -> per-block failures + kv_load_failure_policy="recompute", + ) + try: + env.start_manager() + vllm = env.start_vllm(log_suffix="_p1" if env.hybrid else "") + mbs = env.manager_block_size() + + prompt = make_base_prompts(1, env.hybrid)[0] + suffix = " Now answer: what is 2+2?" + toks = tokenize(vllm.base_url(), prompt) + save_blocks = len(toks) // mbs + self.assertGreaterEqual(save_blocks, 2) + + # ---- Phase 1: save everything. + send_completions(vllm.base_url(), [prompt]) + wait_for_captures(env.capture_dir, "ref", expected=save_blocks, + timeout=180) + self.assertTrue(wait_for_prefix_cached( + env.manager.manager_uri(), env.instance_id, toks, + min_blocks=save_blocks)) + + # ---- Sabotage: delete the tail half of the blocks' files. The + # head blocks stay loadable, so vLLM truncates at the first deleted + # block and the surviving loads remain verifiable. + files = get_block_files(env.manager.manager_uri(), env.instance_id, + toks) + self.assertEqual(len(files), save_blocks) + keep = save_blocks // 2 + for path in files[keep:]: + os.remove(path) + logger.info("deleted %d/%d block files (kept blocks 0..%d)", + save_blocks - keep, save_blocks, keep - 1) + + if env.hybrid: + vllm = env.restart_vllm(log_suffix="_p2") + + # ---- Phase 2: load with holes. The request must return normally. + resp = send_completions(vllm.base_url(), [prompt + suffix])[0] + self.assertTrue(resp["choices"][0]["text"]) + + # The engine must survive and keep serving. + resp2 = send_completions(vllm.base_url(), ["engine alive?"])[0] + self.assertTrue(resp2["choices"][0]["text"]) + + failed_tasks = env.scan_connector_logs(r"load task failed") + self.assertTrue(failed_tasks, "no load failure was logged; the " + "sabotage did not break any loaded block") + + if env.hybrid: + # report_failures=False path: swallowed but logged. + swallowed = env.scan_connector_logs( + r"load failed for \d+/\d+ blocks .*hybrid") + self.assertTrue(swallowed, + "hybrid load failure was not logged") + return + + # Full-attention: vLLM was told about the invalid blocks... + reported = env.scan_connector_logs(r"block_ids_with_load_errors") + self.assertTrue(reported, "failed loads were not reported to vLLM") + + # ...and every surviving head block's loaded KV is bit-exact, + # while any mismatch belongs to a deleted (recomputed) block. + wait_for_captures(env.capture_dir, "loaded", expected=keep, + timeout=180) + report = compare_captures(env.capture_dir, tp_size=1) + hashes = full_block_hashes(toks, mbs) + kept_keys = {("tp0", h) for h in hashes[:keep]} + deleted_keys = {("tp0", h) for h in hashes[keep:]} + failed_keys = {f["key"] for f in report["failures"]} + self.assertFalse( + failed_keys & kept_keys, + f"surviving loaded blocks mismatched: {failed_keys & kept_keys}") + self.assertTrue( + failed_keys <= deleted_keys, + f"mismatches outside the deleted blocks: " + f"{failed_keys - deleted_keys}") + matched_kept = kept_keys & set(report["matched_keys"]) + self.assertEqual( + len(matched_kept), keep, + f"only {len(matched_kept)}/{keep} surviving blocks verified") + finally: + env.stop() + + +if __name__ == "__main__": + unittest.main() diff --git a/integration_test/vllm_e2e/test_multi_turn.py b/integration_test/vllm_e2e/test_multi_turn.py new file mode 100644 index 000000000..5912b0895 --- /dev/null +++ b/integration_test/vllm_e2e/test_multi_turn.py @@ -0,0 +1,112 @@ +"""test_multi_turn: decode-time incremental save feeds the next turn's hit. + +Turn 1 sends a prompt and generates enough output tokens (max_tokens crossing +at least one manager block boundary) that the save threshold in +``build_connector_meta`` fires again during decode: blocks composed of +generated tokens are saved incrementally. Turn 2 sends prompt + turn-1 output +as its prompt (a real multi-turn conversation) and must externally match +*more* blocks than the turn-1 prompt alone covers -- proving decode-produced +blocks were saved -- and their loaded KV must verify against the references +captured during decode. + +Full-attention (mbs=16): three decode blocks, same server both turns (prefix +caching off, the external hit is directly observable). +Hybrid (mbs=528): one decode block (528+ generated tokens); the server is +restarted before turn 2 because prefix caching must stay on for hybrid models +and would otherwise mask the external hit with a local one. +""" + +import logging +import unittest + +from e2e_lib import ( + ScenarioEnv, assert_report_ok, compare_captures, send_completions, + shared_token_prefix_len, tokenize, wait_for_captures, + wait_for_prefix_cached, +) + +logger = logging.getLogger("vllm_e2e") + + +class TestMultiTurn(unittest.TestCase): + def test_multi_turn(self): + env = ScenarioEnv("multi_turn") + try: + env.start_manager() + vllm = env.start_vllm(log_suffix="_t1" if env.hybrid else "") + mbs = env.manager_block_size() + + turn1_prompt = ("A short story request. Please write a long, " + "detailed story about a robot that learns to paint.") + prompt_tokens = tokenize(vllm.base_url(), turn1_prompt) + prompt_blocks = len(prompt_tokens) // mbs + + # ---- Turn 1: generate output crossing >= 1 manager block + # boundary (3 blocks for full-attn's mbs=16; 1 block for hybrid's + # mbs=528 to keep decode time bounded). + max_tokens = (mbs + 32) if env.hybrid else (mbs * 3 + 5) + resp = send_completions( + vllm.base_url(), [turn1_prompt], max_tokens=max_tokens, + ignore_eos=True, return_token_ids=True)[0] + choice = resp["choices"][0] + output_ids = choice["token_ids"] + self.assertEqual(len(output_ids), max_tokens) + turn1_ids = choice["prompt_token_ids"] + output_ids + # The connector tracks tokens when they are *scheduled as input*; + # the very last sampled token never re-enters a step, so at most + # (len - 1) // mbs blocks can have been committed. + turn1_blocks = (len(turn1_ids) - 1) // mbs + self.assertGreater( + turn1_blocks, prompt_blocks, + "turn 1 output did not cross a manager block boundary") + logger.info("turn1: %d prompt + %d output tokens = %d blocks " + "(prompt alone: %d)", len(choice["prompt_token_ids"]), + len(output_ids), turn1_blocks, prompt_blocks) + + # Decode-produced blocks must be committed: the manager holds the + # full prompt+output prefix, more blocks than the prompt covers. + self.assertTrue(wait_for_prefix_cached( + env.manager.manager_uri(), env.instance_id, turn1_ids, + min_blocks=turn1_blocks)) + wait_for_captures(env.capture_dir, "ref", expected=turn1_blocks, + timeout=180) + + if env.hybrid: + # Prefix caching is on for hybrid; restart so turn 2's hit + # comes from KVCM, not the local prefix cache. + vllm = env.restart_vllm(log_suffix="_t2") + + # ---- Turn 2: conversation continues; prompt embeds turn 1's + # prompt + output as token ids (immune to detokenization drift). + turn2_suffix = tokenize(vllm.base_url(), + " Now summarize the story in one word.") + turn2_ids = turn1_ids + turn2_suffix + shared_blocks = min( + shared_token_prefix_len(turn1_ids, turn2_ids) // mbs, + turn1_blocks) + self.assertGreater(shared_blocks, prompt_blocks, + "turn 2 shares no decode-produced block") + resp2 = send_completions(vllm.base_url(), [turn2_ids])[0] + self.assertTrue(resp2["choices"][0]["text"]) + + # The external hit must cover decode-produced blocks. + matched = [int(g[0]) for g in + env.scan_connector_logs(r"matched (\d+) external tokens")] + best = max(matched, default=0) + self.assertGreater( + best, prompt_blocks * mbs, + f"external hit ({best} tokens) does not exceed the prompt-only " + f"coverage ({prompt_blocks * mbs} tokens): decode-time saves " + f"were not used") + + # And the loaded decode-block KV must verify. + wait_for_captures(env.capture_dir, "loaded", + expected=shared_blocks, timeout=180) + report = compare_captures(env.capture_dir, tp_size=1) + assert_report_ok(report, min_matched=shared_blocks) + finally: + env.stop() + + +if __name__ == "__main__": + unittest.main() diff --git a/integration_test/vllm_e2e/test_mutation.py b/integration_test/vllm_e2e/test_mutation.py new file mode 100644 index 000000000..12c4e1e18 --- /dev/null +++ b/integration_test/vllm_e2e/test_mutation.py @@ -0,0 +1,35 @@ +"""test_mutation: meta-test proving the e2e KV verification is not vacuous. + +Runs the standard basic scenario with ``MutatedConnector`` (defined in +test_connector.py, injected only through the test-side +``kv_connector_module_path``), which shifts every attention slot produced by +``_attn_token_indices`` by one -- a symmetric off-by-one: save gathers token +t's KV from the shifted slot and load scatters it back there, so with +contiguous block tables a transport round trip cancels the bug in the interior +of the loaded range. It cannot cancel at the range boundary: one loaded +token's true slot is never written and keeps stale uninitialized data. The +capture comparison reads the cache through vLLM's own slot mapping and must +observe the divergence; run_e2e(expect_verification_failure=True) asserts the +verification FAILS. If the mutated run verifies clean, the harness is blind +and this test fails. +""" + +import unittest + +from e2e_lib import run_e2e + + +class TestMutation(unittest.TestCase): + def test_mutated_connector_is_caught(self): + run_e2e( + scenario="mutation", + tp_size=1, + num_prompts=1, + preferred_block_size=0, + connector_name="MutatedConnector", + expect_verification_failure=True, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/integration_test/vllm_e2e/test_partial_hit.py b/integration_test/vllm_e2e/test_partial_hit.py new file mode 100644 index 000000000..399bd5eb9 --- /dev/null +++ b/integration_test/vllm_e2e/test_partial_hit.py @@ -0,0 +1,96 @@ +"""test_partial_hit: incremental query + incremental save on a non-zero prefix. + +The standard scenarios always start requests from zero computed blocks, so the +``computed_blocks > 0`` branch of ``get_num_new_matched_tokens`` (manager query +with a non-zero ``block_mask.offset``) and the ``block_mask.offset`` branch of +the save path (manager skipping already-stored blocks) are never exercised. +This scenario forces both, with prefix caching enabled for all model types: + +* Stage 1: save prefix A (fresh prefill). +* Stage 2 (same server): send A+B. vLLM locally hits A, so the connector + queries the manager with offset = locally computed blocks (> 0), then + extends the save; the manager's start_write_cache response skips the + already-stored A blocks via a non-zero block_mask offset. +* Stage 3 (restarted server, local cache cleared): send A+B+C. The connector + externally matches A+B -- proving stage 2's incremental save committed -- + loads it, and the loaded KV is verified against the reference captures. + +Log evidence asserted: a getCacheLocation request carrying a non-zero offset +(requires connector DEBUG logging, enabled via log_level). +""" + +import logging +import unittest + +from e2e_lib import ( + ScenarioEnv, assert_report_ok, compare_captures, make_base_prompts, + send_completions, shared_token_prefix_len, tokenize, wait_for_captures, + wait_for_prefix_cached, +) + +logger = logging.getLogger("vllm_e2e") + + +class TestPartialHit(unittest.TestCase): + def test_partial_hit(self): + env = ScenarioEnv("partial_hit", enable_prefix_caching=True, + log_level="DEBUG") + try: + env.start_manager() + vllm = env.start_vllm(log_suffix="_s12") + mbs = env.manager_block_size() + + base = make_base_prompts(1, env.hybrid)[0] + # Three nested prompts: A < A+B < A+B+C. + prompt_a = base + prompt_ab = base + " Continuation section B. " + " ".join( + f"Extra sentence {j} carries value {j * 13 + 7}." + for j in range(90 if env.hybrid else 30)) + prompt_abc = prompt_ab + " Final question: what is 2+2?" + + toks_a = tokenize(vllm.base_url(), prompt_a) + toks_ab = tokenize(vllm.base_url(), prompt_ab) + toks_abc = tokenize(vllm.base_url(), prompt_abc) + blocks_a = len(toks_a) // mbs + blocks_ab = len(toks_ab) // mbs + shared_abc = shared_token_prefix_len(toks_ab, toks_abc) // mbs + self.assertGreaterEqual(blocks_a, 1) + self.assertGreater(blocks_ab, blocks_a, + "B must add at least one manager block") + logger.info("blocks: A=%d AB=%d shared(AB,ABC)=%d", + blocks_a, blocks_ab, shared_abc) + + # ---- Stage 1: save A. + send_completions(vllm.base_url(), [prompt_a]) + wait_for_captures(env.capture_dir, "ref", expected=blocks_a, timeout=180) + self.assertTrue(wait_for_prefix_cached( + env.manager.manager_uri(), env.instance_id, toks_a, + min_blocks=blocks_a)) + + # ---- Stage 2: A hits the local prefix cache -> incremental + # external query (non-zero offset) + incremental save of B. + send_completions(vllm.base_url(), [prompt_ab]) + self.assertTrue(wait_for_prefix_cached( + env.manager.manager_uri(), env.instance_id, toks_ab, + min_blocks=blocks_ab)) + wait_for_captures(env.capture_dir, "ref", expected=blocks_ab, timeout=180) + + offsets = [int(g[0]) for g in env.scan_connector_logs( + r"get_kvcache_location request:.*'offset': (\d+)")] + self.assertTrue(any(o > 0 for o in offsets), + f"no incremental query with non-zero offset: {offsets}") + + # ---- Stage 3: restart (clear local cache) and load A+B. + vllm = env.restart_vllm(log_suffix="_s3") + send_completions(vllm.base_url(), [prompt_abc]) + wait_for_captures(env.capture_dir, "loaded", + expected=shared_abc, timeout=180) + + report = compare_captures(env.capture_dir, tp_size=1) + assert_report_ok(report, min_matched=shared_abc) + finally: + env.stop() + + +if __name__ == "__main__": + unittest.main()