From ec9c365708dbdc78ef6a748db9e543a8b15ee64d Mon Sep 17 00:00:00 2001 From: Ketor Date: Thu, 17 Sep 2026 15:09:34 +0800 Subject: [PATCH 1/3] fix(vllm): recover hybrid KV loads through request-level outcomes --- docs/CONNECTORS.md | 32 +- integration/vllm/README.md | 15 +- .../patches/vllm-hybrid-invalid-blocks.patch | 62 ---- integration/vllm/src/dfkv_vllm/connector.py | 15 +- integration/vllm/src/dfkv_vllm/data.py | 12 + integration/vllm/src/dfkv_vllm/scheduler.py | 37 +- integration/vllm/src/dfkv_vllm/worker.py | 181 +++++++--- .../vllm/tests/test_hybrid_pool_layout.py | 1 + .../vllm/tests/test_request_load_failures.py | 318 ++++++++++++++++++ .../tests/test_scheduler_full_block_ids.py | 155 +++++++++ .../tests/test_stateful_lookup_admission.py | 12 +- .../vllm/tests/test_worker_lifecycle.py | 2 +- 12 files changed, 694 insertions(+), 148 deletions(-) delete mode 100644 integration/vllm/patches/vllm-hybrid-invalid-blocks.patch create mode 100644 integration/vllm/tests/test_request_load_failures.py diff --git a/docs/CONNECTORS.md b/docs/CONNECTORS.md index c1aefd1..96a8084 100644 --- a/docs/CONNECTORS.md +++ b/docs/CONNECTORS.md @@ -882,19 +882,23 @@ lsmod | grep nvidia_peermem #### 混合模型的 KV 加载故障恢复 -- vLLM 上游 `_handle_invalid_blocks` 的 `(req_block_ids,) = ...` 解包在多 KV-group - 模型上必然抛 `ValueError`,导致 EngineCore 死亡。随包提供 - [混合 invalid-blocks 修复补丁](../integration/vllm/patches/vllm-hybrid-invalid-blocks.patch): - hybrid 请求按 per-group 外层 spec(AttentionSpec 乘 DCP)计算受影响前缀, - 命中失败时使所有参与 group 的相关 block hash 失效(仅 hash,不释放 DMA 中 - buffer),`skip_reading_prefix_cache=True` 后复用 `_preempt_request` 回到 - waiting 队列本地重算;async 请求在 finished_recving 之后由既有 - `_update_waiting_for_remote_kv` 释放。单 group 请求保留原最长有效前缀与 - 共享 block 优化。fail 策略仅上报受影响请求与 eviction 集合,不做重放。 - dfkv 调度器同步配合:重放请求查询前丢弃残余 lookup/load_spec 并返回 - `(0, False)`,不会反复撞同一外部缓存。补丁针对实测镜像 - `vllm/vllm-openai:glm53-flash` 的 `vllm/v1/core/sched/scheduler.py`,应用前先 - `git apply --check`;不可盲目覆盖其它引擎版本。 +- 混合/多 group 和单 group 非 full-attention 布局使用 vLLM 原生请求级失败 + 协议,要求引擎提供 `KVConnectorTransferResults.failed_recving`、 + `KVConnectorOutput.failed_recving` 的完整透传,以及对 + `WAITING_FOR_REMOTE_KVS` 请求的调度器恢复处理。请升级到具备完整 API 的 + 引擎;不再提供或要求外置 invalid-blocks 调度器补丁,也不按未经验证的 + 版本号推定支持。 +- 这些布局按 request ID 报错,`invalid_block_ids` 留空;只有原生 I/O 和 GPU + 写入都经过终止 fence 后,才同时报告 `failed_recving` 与 `finished_recving`。 + `kv_load_failure_policy="recompute"` 时,引擎释放失败分配并重新本地准入; + dfkv 丢弃旧 lookup、load spec 和 tracker,按新分配及实际计算前缀重建 SAVE + 状态。失败请求在本次生命周期内(包括后续抢占)不再查询远端命中,避免同一 + 广告命中不断触发加载—失败循环;最终完成时清理此状态。`"fail"` 策略由引擎 + 终止受影响请求。单 group full-attention 保留既有 block-level 错误恢复。 +- `load_async=False` 只选择串行模型线程 I/O,不取消上述请求的 + `WAITING_FOR_REMOTE_KVS` 准入。连接器在 GET 前等待先前 GPU 工作完成, + GET 后再完成写入 fence;按 runner 调用顺序可能处于前一个 forward 之后, + 但不与 GPU compute 重叠。失败请求本身在恢复完成前不执行 forward。 ### 3.2 验证 @@ -940,7 +944,7 @@ namespace/key 不一致是预期 cold miss。**空环 / MDS 不可达**可直接 | `batch_concurrency` | `8` | **大池可调高到 ≈ 节点数** | 跨节点 fan-out,**真正的吞吐杠杆**(depth 是平的) | | `rail_affinity` | `False` | 多 rank、多 rail 生产设 `true` | 按 vLLM world-group per-host `local_rank` 选择 primary;在 native client 创建前设置每进程独立 rail 环境 | | `rail_affinity_fallbacks` | `1` | `1` | 相邻有序 fallback 数;`0`=严格单 rail,超出可用 rail 数时自动收敛 | -| `load_async` | `True` | 普通 attention 保持 True;hybrid recurrent 模型设 `False` | `False` 在 forward 前同步完成 load,避免 recurrent-state compute 与远端 GPU 写重叠 | +| `load_async` | `True` | 普通 attention 保持 True;hybrid recurrent 模型设 `False` | `False` 在模型线程串行执行 load,GET 前后均等待 GPU fence,可能位于前一个 forward 之后;hybrid/多 group 和非 full-attention 请求仍进入 `WAITING_FOR_REMOTE_KVS`,不会消费失败 KV。单 group full-attention 保留同步准入。 | | `transfer_queue_capacity` | `256` | 保持默认,按压测调 | 每个 worker、每个方向的排队上限(`1..65536`)。满队列时非阻塞拒绝新任务:save 立即释放 finish/free fence,load 标记失败并重算;非法值启动即失败。 | | `recv_workers` | `1` | 从 `1` 起压测 | 共享有界 receive queue 的 GET worker 数(`1..32`);仅在 queue wait 持续升高且后端仍有余量时增加。 | | `load_window_keys` | `0`(关闭) | 长上下文 replicated-MLA 按压测设置 | 单次 native GET window 的最大 key 数(`0..65536`)。窗口结果应能放入 `DFKV_NODE_DEDUP_GPU_ARENA_MB`,且在 `DFKV_NODE_DEDUP_WAIT_MS` 内完成,避免 follower rank 超时后重复读取同一批 KV。 | diff --git a/integration/vllm/README.md b/integration/vllm/README.md index c7be46b..4ec6f85 100644 --- a/integration/vllm/README.md +++ b/integration/vllm/README.md @@ -41,6 +41,19 @@ device pointers. Set `DFKV_RDMA=1` in every engine process. Construction rejects and closes any native handle whose reported transport is not `rdma`; there is no TCP or host-bounce fallback for this connector. +Hybrid/multi-group caches and single-group non-full-attention caches require +vLLM's native request-level load-failure protocol: +`KVConnectorTransferResults.failed_recving`, its propagation through +`KVConnectorOutput.failed_recving`, and scheduler recovery of failed +`WAITING_FOR_REMOTE_KVS` requests. Upgrade to an engine exposing that complete +API; no out-of-tree invalid-block scheduler patch is supported or required. +The connector reports failed request IDs together with receive completion only +after native I/O and GPU writes are fenced, leaving block-ID errors empty for +these layouts. With `kv_load_failure_policy="recompute"`, the engine releases the +failed allocation and retries locally; dfkv bypasses the failed remote source +for the remainder of that request, including subsequent preemptions. The engine's +`"fail"` policy terminates the affected request instead. + ## Environment variables (engine process) Read by `libdfkv.so` (the C client) and the connector, so set them in **every** @@ -87,7 +100,7 @@ LMCache connector access logs, so one setting covers every integration. Format: | `batch_concurrency` | `0`=auto | client fan-out for batch ops; the real throughput lever (depth is flat). Auto = `min(max(nodes, 8), 32)`: 8-way parallel on single-node, one-per-node on multi-node. Set >0 to pin a fixed value. | | `rail_affinity` | `False` | Bind each vLLM worker process to a primary rail selected by world-group local rank; requires an ordered multi-rail `DFKV_RDMA_DEV`. | | `rail_affinity_fallbacks` | `1` | Number of ordered neighboring fallback rails when affinity is enabled. `0` keeps strict one-rank/one-rail; values above the available rail count are bounded. | -| `load_async` | `True` | `True` returns `WAITING_FOR_REMOTE_KVS` and overlaps GPUDirect loads with unrelated model work. `False` performs each requested load synchronously in `start_load_kv`, before the forward pass. Use `False` for hybrid state-cache models when the engine cannot guarantee that remote writes target blocks disjoint from concurrent compute. | +| `load_async` | `True` | Controls I/O execution, not hybrid admission. `True` overlaps GPUDirect loads with unrelated model work. `False` executes serialized loads on the model thread, fences preceding GPU work before GET, and fences completion before returning; depending on the runner, this may occur after the preceding forward. Hybrid/multi-group and non-full-attention requests still enter `WAITING_FOR_REMOTE_KVS` in either mode, so failed KV is never consumed by their forward pass. Single-group full-attention `False` loads retain synchronous admission. Use `False` for hybrid state-cache models when remote writes cannot safely overlap compute. | | `transfer_queue_capacity` | `256` | Maximum queued requests in each direction (`1..65536`). All receive workers consume one shared receive queue of this capacity; capacity is not multiplied by `recv_workers`. Submission is non-blocking: a full queue rejects new saves as completed (releasing finish/free fences) and rejects new loads as load errors (forcing recompute), so overload cannot grow memory or pin blocks indefinitely. Invalid or out-of-range values abort connector construction. | | `recv_workers` | `1` | Receive/load worker count (`1..32`). Workers consume the shared bounded receive queue and may execute independent native GETs concurrently. Invalid, boolean, or out-of-range values abort connector construction. | | `load_window_keys` | `0` (disabled) | Maximum keys per native GET window (`0..65536`). Use a value whose worst-case result bytes fit inside the node-dedup GPU arena. Windowing lets follower ranks consume published results before the dedup wait deadline instead of re-fetching a large replicated-MLA batch. | diff --git a/integration/vllm/patches/vllm-hybrid-invalid-blocks.patch b/integration/vllm/patches/vllm-hybrid-invalid-blocks.patch deleted file mode 100644 index 924eb76..0000000 --- a/integration/vllm/patches/vllm-hybrid-invalid-blocks.patch +++ /dev/null @@ -1,62 +0,0 @@ ---- a/vllm/v1/core/sched/scheduler.py -+++ b/vllm/v1/core/sched/scheduler.py -@@ -2839,7 +2839,58 @@ - req_num_computed_tokens = ( - request.num_computed_tokens - num_scheduled_tokens.get(req_id, 0) - ) -- # TODO (davidb): add support for hybrid memory allocator -+ if len(group_block_ids) > 1: -+ req_num_computed_tokens = max(0, req_num_computed_tokens) -+ dependent_groups = [] -+ for block_ids, group in zip( -+ group_block_ids, -+ self.kv_cache_config.kv_cache_groups, -+ strict=True, -+ ): -+ spec = group.kv_cache_spec -+ if not spec.participates_in_prefix_caching: -+ continue -+ # Match the engine's outer-spec DCP coverage rule. -+ block_size = spec.block_size -+ if isinstance(spec, AttentionSpec): -+ block_size *= self.dcp_world_size -+ dependent_groups.append(block_ids) -+ num_computed_blocks = ( -+ req_num_computed_tokens + block_size - 1 -+ ) // block_size -+ if any( -+ block_id in invalid_block_ids -+ for _, block_id in zip(range(num_computed_blocks), block_ids) -+ ): -+ is_affected = True -+ -+ if not is_affected: -+ continue -+ affected_req_ids.add(req_id) -+ total_affected_tokens += req_num_computed_tokens -+ # A failed group invalidates all dependent group states, not -+ # just the failed group's suffix. evict_blocks removes hashes -+ # only: it must not return DMA-owned buffers to the pool. -+ dependent_block_ids = { -+ block_id for block_ids in dependent_groups for block_id in block_ids -+ } -+ if evict_blocks: -+ blocks_to_evict.update(dependent_block_ids) -+ if self.recompute_kv_load_failures: -+ self.kv_cache_manager.evict_blocks(dependent_block_ids) -+ request.skip_reading_prefix_cache = True -+ if request.status == RequestStatus.RUNNING: -+ self.running.remove(request) -+ self._preempt_request( -+ request, time.monotonic(), drop_stale_output=True -+ ) -+ else: -+ assert request.status == RequestStatus.WAITING_FOR_REMOTE_KVS -+ # Keep async receive buffers until finished_recving; -+ # _update_waiting_for_remote_kv then frees all groups. -+ request.num_computed_tokens = 0 -+ continue -+ - (req_block_ids,) = group_block_ids - - req_num_computed_blocks = ( diff --git a/integration/vllm/src/dfkv_vllm/connector.py b/integration/vllm/src/dfkv_vllm/connector.py index 6c897ad..dc8f619 100644 --- a/integration/vllm/src/dfkv_vllm/connector.py +++ b/integration/vllm/src/dfkv_vllm/connector.py @@ -27,6 +27,7 @@ KVConnectorBase_V1, KVConnectorMetadata, KVConnectorRole, + KVConnectorTransferResults, SupportsHMA, ) from vllm.distributed.kv_transfer.kv_connector.v1.metrics import ( @@ -244,6 +245,8 @@ def reset_cache(self) -> bool | None: self._finish_call() def update_connector_output(self, connector_output: KVConnectorOutput): + assert self.connector_scheduler is not None + self.connector_scheduler.update_connector_output(connector_output) kv_cache_events = connector_output.kv_cache_events if not kv_cache_events or not isinstance( kv_cache_events, DfkvStoreKVEvents @@ -319,8 +322,8 @@ def handle_preemptions( def start_load_kv(self, forward_context: ForwardContext, **kwargs: Any) -> None: self._begin_call() try: - # Loads are issued in get_finished() for compute overlap. Synchronous - # loads required by this step are submitted here before forward. + # Async loads are issued during result collection. Inline loads + # fence prior kernels before touching destination blocks. assert self.connector_worker is not None metadata = self._get_connector_metadata() assert isinstance(metadata, DfkvStoreConnectorMetadata) @@ -353,20 +356,20 @@ def save_kv_layer( def wait_for_save(self): self._begin_call() try: - # get_finished submits stores and fences mutable/windowed sources. + # Result collection submits stores and fences mutable/windowed sources. return finally: self._finish_call() - def get_finished( + def get_transfer_results( self, finished_req_ids: set[str] - ) -> tuple[set[str] | None, set[str] | None]: + ) -> KVConnectorTransferResults: self._begin_call() try: assert self.connector_worker is not None metadata = self._get_connector_metadata() assert isinstance(metadata, DfkvStoreConnectorMetadata) - return self.connector_worker.get_finished(finished_req_ids, metadata) + return self.connector_worker.get_transfer_results(finished_req_ids, metadata) finally: self._finish_call() diff --git a/integration/vllm/src/dfkv_vllm/data.py b/integration/vllm/src/dfkv_vllm/data.py index b36f2aa..641d3f4 100644 --- a/integration/vllm/src/dfkv_vllm/data.py +++ b/integration/vllm/src/dfkv_vllm/data.py @@ -31,6 +31,18 @@ # they must cold-miss rather than enter the current restore path. VLLM_RAW_LAYOUT = b"vllm-multiwr-v4" + +def requires_request_level_loads(kv_cache_config) -> bool: + """Use request identities when physical groups cannot share a block grid.""" + from vllm.v1.kv_cache_interface import FullAttentionSpec + + groups = getattr( + kv_cache_config, "transfer_groups", kv_cache_config.kv_cache_groups + ) + return len(groups) > 1 or any( + not isinstance(group.kv_cache_spec, FullAttentionSpec) for group in groups + ) + def key_diagnostic_label(key: bytes) -> str: """Return the standard non-reversible diagnostic label for a store key.""" try: diff --git a/integration/vllm/src/dfkv_vllm/scheduler.py b/integration/vllm/src/dfkv_vllm/scheduler.py index 3a65eb9..1899a0c 100644 --- a/integration/vllm/src/dfkv_vllm/scheduler.py +++ b/integration/vllm/src/dfkv_vllm/scheduler.py @@ -17,6 +17,7 @@ from vllm.v1.core.kv_cache_utils import resolve_kv_cache_block_sizes from vllm.v1.core.sched.output import NewRequestData, SchedulerOutput from vllm.v1.kv_cache_interface import KVCacheConfig +from vllm.v1.outputs import KVConnectorOutput from vllm.v1.request import Request from ._determinism import ensure_deterministic_block_hashing @@ -25,6 +26,7 @@ DfkvStoreConnectorMetadata, ReqMeta, RequestTracker, + requires_request_level_loads, ) from .worker import ( LookupKeyClient, @@ -66,6 +68,7 @@ def __init__( self.load_async = extra_config.get("load_async", True) if not isinstance(self.load_async, bool): raise ValueError("dfkv connector: load_async must be a boolean") + self.request_level_loads = requires_request_level_loads(kv_cache_config) self.lookup_async = extra_config.get("lookup_async", False) self.client = LookupKeyClient(vllm_config) self._closed = False @@ -82,6 +85,7 @@ def __init__( self._unfinished_requests: dict[str, tuple[Request, tuple[list[int], ...]]] = {} self._unfinished_request_ids: set[str] = set() self._allocated_req_ids: set[str] = set() + self._failed_load_req_ids: set[str] = set() def get_num_new_matched_tokens( self, @@ -93,7 +97,10 @@ def get_num_new_matched_tokens( Returns ``(None, False)`` while an asynchronous lookup is pending so vLLM retries the request on a later scheduler step. """ - if getattr(request, "skip_reading_prefix_cache", False): + if ( + request.request_id in self._failed_load_req_ids + or getattr(request, "skip_reading_prefix_cache", False) + ): self.client.discard(request.request_id) self.load_specs.pop(request.request_id, None) return 0, False @@ -141,7 +148,9 @@ def get_num_new_matched_tokens( can_load=False, ) - return need_to_allocate, self.load_async + # Request-level failure recovery only applies to waiting requests. + # Parking admission is independent of whether the worker overlaps I/O. + return need_to_allocate, self.load_async or self.request_level_loads def update_state_after_alloc( self, @@ -177,6 +186,27 @@ def update_state_after_alloc( self.load_specs[request.request_id].can_load = True + def update_connector_output(self, output: KVConnectorOutput) -> None: + """Prevent a failed remote source from being re-admitted on retry.""" + if not self.request_level_loads: + return + for req_id in output.failed_recving or (): + # A receive can finish after an aborted request's final metadata. + # Such an output must not recreate scheduler-side request state. + if req_id not in self._unfinished_request_ids: + continue + self._failed_load_req_ids.add(req_id) + self.client.discard(req_id) + self.load_specs.pop(req_id, None) + self._request_trackers.pop(req_id, None) + self._unfinished_requests.pop(req_id, None) + self._allocated_req_ids.discard(req_id) + # Upstream frees failed-load blocks and re-admits without emitting + # a new preemption. MRV1 may use scheduled_cached_reqs on that retry; + # rebuild its tracker from the new allocation and computed prefix, + # never from the failed load's advertised token count/block table. + self._preempted_req_ids.add(req_id) + def build_connector_meta( self, scheduler_output: SchedulerOutput ) -> KVConnectorMetadata: @@ -190,6 +220,7 @@ def build_connector_meta( self._unfinished_requests.pop(finished_req_id, None) self._unfinished_request_ids.discard(finished_req_id) self._preempted_req_ids.discard(finished_req_id) + self._failed_load_req_ids.discard(finished_req_id) preempted_ids = scheduler_output.preempted_req_ids or set() self._preempted_req_ids.update(preempted_ids) @@ -375,6 +406,8 @@ def request_finished( ) -> tuple[bool, dict[str, Any] | None]: """Determine whether to delay freeing blocks for async save.""" self.client.discard(request.request_id) + self._failed_load_req_ids.discard(request.request_id) + self.load_specs.pop(request.request_id, None) if self.kv_role == "kv_consumer": return False, None tracker = self._request_trackers.get(request.request_id) diff --git a/integration/vllm/src/dfkv_vllm/worker.py b/integration/vllm/src/dfkv_vllm/worker.py index ac1c788..9574a62 100644 --- a/integration/vllm/src/dfkv_vllm/worker.py +++ b/integration/vllm/src/dfkv_vllm/worker.py @@ -48,6 +48,9 @@ get_tensor_model_parallel_world_size, ) from vllm.distributed.kv_events import BlockStored +from vllm.distributed.kv_transfer.kv_connector.v1.base import ( + KVConnectorTransferResults, +) from vllm.logger import init_logger from vllm.utils.network_utils import make_zmq_socket from vllm.v1.core import kv_cache_utils @@ -81,6 +84,7 @@ PoolKey, ReqMeta, key_diagnostic_label, + requires_request_level_loads, split_block_contiguous_runs, ) from .dfkv_client import DfkvDeviceClient, SgDescriptorBatch @@ -299,7 +303,7 @@ def _logical_block_ids( # Each direction owns one queue. ReqMeta objects retain block/hash lists and # CUDA-event references, so an unbounded queue can grow host memory and keep # scheduler-owned GPU blocks pinned indefinitely when the native client slows. -# Reject-new is deliberately non-blocking: blocking get_finished() on queue +# Reject-new is deliberately non-blocking: blocking result collection on queue # capacity would prevent vLLM from observing completions and freeing blocks. DEFAULT_TRANSFER_QUEUE_CAPACITY = 256 MAX_TRANSFER_QUEUE_CAPACITY = 65536 @@ -424,6 +428,7 @@ class _ReceiveRequestState: phase: str = "queued" cancel_requested: bool = False fail_closed_on_cancel: bool = False + failed: bool = False completion: threading.Event = dataclasses.field(default_factory=threading.Event) @@ -963,6 +968,7 @@ def __init__( load_window_keys: int = DEFAULT_LOAD_WINDOW_KEYS, load_window_min_keys: int = DEFAULT_LOAD_WINDOW_MIN_KEYS, record_pool_sample: Callable[[str, int], None] | None = None, + request_level_loads: bool = False, ): super().__init__( client, @@ -987,6 +993,8 @@ def __init__( load_window_min_keys ) self.client_provider = client_provider + self.request_level_loads = request_level_loads + self._failed_requests: set[str] = set() self._invalid_block_ids_lock = threading.Lock() self._invalid_block_ids: set[int] = set() self.coord = coord @@ -1057,7 +1065,6 @@ def add_request(self, request: ReqMeta) -> bool: return False self._request_states[request.req_id] = state self._terminalize_locked(state, failed=True) - self._request_states.pop(request.req_id, None) logger.warning( "%s rejected request %s: transfer queue %s (capacity=%d)", self.name, @@ -1067,15 +1074,22 @@ def add_request(self, request: ReqMeta) -> bool: ) return False - def get_and_clear_finished_requests(self) -> set[str]: + def get_and_clear_receive_results(self) -> tuple[set[str], set[str]]: + """Drain terminal completions and failures in one atomic snapshot.""" with self._request_states_lock: with self.done_task_lock: - finished = self.finished_requests.copy() - self.finished_requests.clear() + finished = self.finished_requests + failed = self._failed_requests + self.finished_requests = set() + self._failed_requests = set() for req_id in finished: state = self._request_states.get(req_id) if state is not None and state.phase == "terminal": self._request_states.pop(req_id, None) + return finished, failed + + def get_and_clear_finished_requests(self) -> set[str]: + finished, _ = self.get_and_clear_receive_results() return finished def cancel_requests( @@ -1116,13 +1130,16 @@ def _terminalize_locked( req_id = req.req_id if req is not None else None state.phase = "terminal" state.request = None - if failed and state.block_ids: + failed = failed or state.failed + if failed and not self.request_level_loads and state.block_ids: with self._invalid_block_ids_lock: self._invalid_block_ids.update(state.block_ids) state.block_ids = () if req_id is not None: with self.done_task_lock: self.finished_requests.add(req_id) + if failed and self.request_level_loads: + self._failed_requests.add(req_id) state.completion.set() return True @@ -1206,6 +1223,15 @@ def stop(self, *, cancel_pending: bool = True) -> None: for state in self._request_states.values(): if state.phase == "active": state.cancel_requested = True + # Inline model-thread loads share the ownership fence but are not + # among worker_threads; shutdown must also join those native calls. + with self._request_states_lock: + active_completions = [ + state.completion for state in self._request_states.values() + if state.phase == "active" + ] + for completion in active_completions: + completion.wait() if self.ident is None: while True: try: @@ -1240,19 +1266,60 @@ def stop(self, *, cancel_pending: bool = True) -> None: close = stop - def _add_load_error_block_ids(self, block_ids: list[int]) -> None: - with self._invalid_block_ids_lock: - self._invalid_block_ids.update(block_ids) + def _record_load_failure(self, req_id: str, block_ids: list[int]) -> None: + if self.request_level_loads: + with self._request_states_lock: + self._request_states[req_id].failed = True + else: + with self._invalid_block_ids_lock: + self._invalid_block_ids.update(block_ids) def get_and_clear_block_ids_with_load_errors(self) -> set[int]: + if self.request_level_loads: + return set() with self._invalid_block_ids_lock: invalid_block_ids = self._invalid_block_ids.copy() self._invalid_block_ids.clear() return invalid_block_ids def load_request_sync(self, request: ReqMeta) -> None: - """Load one request on the model thread before its forward pass.""" - self._handle_request(request) + """Load on the model thread; parked requests report outcomes later.""" + if not self.request_level_loads: + self._handle_request(request) + return + with self._stop_lock: + with self._request_states_lock: + if request.req_id in self._request_states: + return + state = _ReceiveRequestState( + request=request, + block_ids=tuple( + block for group in request.block_ids for block in group + ), + phase="active", + ) + self._request_states[request.req_id] = state + if not self._accepting: + self._terminalize_locked(state, failed=True) + return + failed = False + try: + if self._cuda_device is not None: + # A runner may defer this hook until after forward for parked + # loads. Finish outstanding kernels before modifying the pool. + torch.cuda.synchronize(self._cuda_device) + self._handle_request(request) + except Exception: + failed = True + raise + finally: + with self._request_states_lock: + self._terminalize_locked( + state, + failed=failed or ( + state.cancel_requested and state.fail_closed_on_cancel + ), + ) def _handle_request(self, req_meta: ReqMeta): @@ -1360,7 +1427,7 @@ def _handle_request(self, req_meta: ReqMeta): f"keys={len(rotated_keys)} hits={len(hits)} lens={len(lens)}" ) except Exception as e: - self._add_load_error_block_ids(rotated_block_ids) + self._record_load_failure(req_id, rotated_block_ids) self._record_operation( "load_get", load_get_start, @@ -1380,22 +1447,20 @@ def _handle_request(self, req_meta: ReqMeta): e, ) return + finally: + # Fence even a failed/partial native batch before publishing + # its terminal outcome and allowing destination block reuse. + if ( + os.environ.get("DFKV_GPU_LOAD_FENCE", "1") == "1" + and self._cuda_device is not None + ): + torch.cuda.synchronize(self._cuda_device) failed_indices = [ i for i, (hit, got_len) in enumerate(zip(hits, lens, strict=True)) if hit != 1 or got_len != chunk_totals[i] ] - # GPUDirect RDMA ordering fence: the CQ completion proves - # transmission, not arrival of the BAR writes in device memory; - # drivers reject CU_POINTER_ATTRIBUTE_SYNC_MEMOPS on VMM pools, so - # synchronize the device before the scheduler launches kernels over - # the loaded blocks. DFKV_GPU_LOAD_FENCE=0 disables. - if ( - os.environ.get("DFKV_GPU_LOAD_FENCE", "1") == "1" - and self._cuda_device is not None - ): - torch.cuda.synchronize(self._cuda_device) failed_block_ids = [rotated_block_ids[i] for i in failed_indices] self._record_operation( "load_get", @@ -1414,7 +1479,7 @@ def _handle_request(self, req_meta: ReqMeta): len(failed_block_ids), ) if failed_block_ids: - self._add_load_error_block_ids(failed_block_ids) + self._record_load_failure(req_id, failed_block_ids) if logger.isEnabledFor(logging.WARNING): failed_detail = [ ( @@ -1438,12 +1503,9 @@ def _handle_request(self, req_meta: ReqMeta): # Any unexpected failure in the load path -> recompute this # request's blocks (never hang vLLM's WAITING_FOR_REMOTE_KVS). logger.error("dfkv recv thread failed for req %s: %s", req_id, e) - try: - self._add_load_error_block_ids( - [b for ids in req_meta.block_ids for b in ids] - ) - except Exception: - pass + self._record_load_failure( + req_id, [b for ids in req_meta.block_ids for b in ids] + ) def _cancel_request(self, req_meta: Any) -> None: self.cancel_requests( (req_meta.req_id,), wait=False, fail_closed=True @@ -1552,9 +1614,10 @@ def __init__( self.load_async = extra.get("load_async", True) if not isinstance(self.load_async, bool): raise ValueError("dfkv connector: load_async must be a boolean") + self.request_level_loads = requires_request_level_loads(kv_cache_config) logger.info( "dfkv load mode: %s", - "async-overlap" if self.load_async else "synchronous-before-forward", + "async-overlap" if self.load_async else "serialized-model-thread", ) self.cache_config = vllm_config.cache_config self.block_size, self.hash_block_size = resolve_kv_cache_block_sizes( @@ -1865,7 +1928,7 @@ def register_cross_layers_kv_caches(self, kv_cache: torch.Tensor) -> None: def _ensure_client_for_load(self) -> Any: """Lazily un-elide: create the dfkv client on an elided producer rank the first time a real load reaches it (cross-instance prefix reuse — - get_finished has no role gate). Keeps phase 2a's connection savings for + get_transfer_results has no role gate). Keeps phase 2a's connection savings for the common P-instance case while never trading a whole-span recompute for them. Thread-safe; returns None (load misses, vLLM recomputes) if creation fails or the rank was never elided-with-kwargs.""" @@ -2074,6 +2137,7 @@ def _repr_tensor(v: torch.Tensor | list[torch.Tensor]) -> torch.Tensor: "load_window_min_keys", DEFAULT_LOAD_WINDOW_MIN_KEYS, ), + request_level_loads=self.request_level_loads, ) self.kv_recv_thread.start() ready_event_recving.wait() @@ -2112,9 +2176,10 @@ def start_load_kv( self, metadata: DfkvStoreConnectorMetadata, ): - """Perform synchronous loads before forward.""" - if self.load_async: + """Serialize native loads on the model thread, never a pool thread.""" + if self.load_async or getattr(self, "_last_load_metadata", None) is metadata: return + self._last_load_metadata = metadata assert self.kv_recv_thread is not None for request in metadata.requests: load_spec = request.load_spec @@ -2124,12 +2189,12 @@ def start_load_kv( self.kv_recv_thread.load_request_sync(request) - def get_finished( + def get_transfer_results( self, finished_req_ids: set[str], meta: DfkvStoreConnectorMetadata, - ) -> tuple[set[str], set[str]]: - """Submit post-forward I/O and collect completed request IDs. + ) -> KVConnectorTransferResults: + """Submit post-forward I/O and atomically collect receive outcomes. Mutable and windowed stores finish before the next model step can overwrite their sources; full-attention stores retain async overlap. @@ -2142,8 +2207,31 @@ def get_finished( wait=True, fail_closed=False, ) + if getattr(self, "_last_transfer_metadata", None) is not meta: + self._last_transfer_metadata = meta + self._submit_transfers(meta) + + done_sending = ( + self._get_and_clear_finished_sending(finished_req_ids) + if self.kv_role in ["kv_producer", "kv_both"] + else set() + ) + done_recving, failed_recving = set(), set() + if self.kv_recv_thread is not None and ( + self.load_async or self.request_level_loads + ): + done_recving, failed_recving = ( + self.kv_recv_thread.get_and_clear_receive_results() + ) + return KVConnectorTransferResults( + finished_sending=done_sending, + finished_recving=done_recving, + failed_recving=failed_recving, + ) + + def _submit_transfers(self, meta: DfkvStoreConnectorMetadata) -> None: # Async mode overlaps loads with unrelated model work. Synchronous mode - # already completed them in start_load_kv, before this forward pass. + # executes them inline in start_load_kv, never in the receive pool. if self.load_async: for request in meta.requests: load_spec = request.load_spec @@ -2176,25 +2264,6 @@ def get_finished( # Preserve those source bytes until the native PUT completes. self.kv_send_thread.request_queue.join() - # Check completion of previously queued transfers - done_sending = ( - self._get_and_clear_finished_sending(finished_req_ids) - if self.kv_role in ["kv_producer", "kv_both"] - else set() - ) - - done_recving = ( - self.kv_recv_thread.get_and_clear_finished_requests() - if self.load_async and self.kv_recv_thread is not None - else set() - ) - - if done_sending or done_recving: - logger.debug( - "dfkv get_finished: done_recving=%s done_sending=%s tp=%d", - done_recving, done_sending, self.tp_rank, - ) - return done_sending, done_recving def get_block_ids_with_load_errors(self) -> set[int]: if self.kv_recv_thread is None: diff --git a/integration/vllm/tests/test_hybrid_pool_layout.py b/integration/vllm/tests/test_hybrid_pool_layout.py index 030de1b..f1f48db 100644 --- a/integration/vllm/tests/test_hybrid_pool_layout.py +++ b/integration/vllm/tests/test_hybrid_pool_layout.py @@ -186,6 +186,7 @@ def start(self): ) worker = DfkvStoreWorker.__new__(DfkvStoreWorker) worker._kv_cache_groups = groups + worker.request_level_loads = False worker.token_dbs = [db] worker.cache_config = SimpleNamespace(num_gpu_blocks=8) worker._kv_pool_regions = [] diff --git a/integration/vllm/tests/test_request_load_failures.py b/integration/vllm/tests/test_request_load_failures.py new file mode 100644 index 0000000..7296988 --- /dev/null +++ b/integration/vllm/tests/test_request_load_failures.py @@ -0,0 +1,318 @@ +"""Request-level receives publish one fenced completion/failure snapshot.""" + +import ctypes +import threading +from types import SimpleNamespace + +import pytest + +pytest.importorskip("vllm") + +from vllm.v1.core.kv_cache_utils import BlockHash + +from dfkv_vllm.connector import DfkvStoreConnector +from dfkv_vllm.data import ( + ChunkedTokenDatabase, + DfkvStoreConnectorMetadata, + KeyMetadata, + LoadSpec, + ReqMeta, +) +from dfkv_vllm.worker import DfkvStoreWorker, KVCacheStoreRecvingThread + + +BLOCK = 16 + + +def request(req_id="load", block_id=1): + return ReqMeta( + req_id=req_id, token_len_chunk=BLOCK, block_ids=([block_id],), + block_hashes=[BlockHash(b"x" * 32)], + load_spec=LoadSpec(0, BLOCK, True, token_len=BLOCK), + ) + + +class MemoryClient: + def __init__(self, outcome="ok"): + self.outcome = outcome + self.calls = 0 + self.thread_ids = [] + + def batch_get_auto_sg(self, keys, pointers, capacities): + self.calls += 1 + self.thread_ids.append(threading.get_ident()) + for ptrs, caps in zip(pointers, capacities, strict=True): + for ptr, cap in zip(ptrs, caps, strict=True): + ctypes.memset(ptr, 17, cap) + if self.outcome == "native": + raise RuntimeError("native GET failed after a partial write") + if self.outcome == "incomplete": + return [], [] + if self.outcome == "miss": + return [False] * len(keys), [0] * len(keys) + length = {"short": BLOCK - 1, "oversized": BLOCK + 1}.get( + self.outcome, BLOCK, + ) + return [True] * len(keys), [length] * len(keys) + + +@pytest.fixture +def make_receiver(monkeypatch): + # The fixture tests native-call ownership with CPU pointers. GPU ordering + # is exercised separately by explicitly controlled owner-device fences. + monkeypatch.setattr("torch.cuda.is_available", lambda: False) + receivers = [] + pools = [] + + def make(client=None, *, request_level=True, workers=1, capacity=4): + pool = ctypes.create_string_buffer(4 * BLOCK) + pools.append(pool) + metadata = KeyMetadata( + model_name="request-recovery", dp_size=1, dp_rank=-1, + tp_size=1, tp_rank=0, pcp_size=1, pcp_rank=0, + dcp_size=1, dcp_rank=0, pp_size=1, pp_rank=0, + ) + database = ChunkedTokenDatabase(metadata, BLOCK, hash_block_size=BLOCK) + database.set_seg_layout([(ctypes.addressof(pool), BLOCK, BLOCK)]) + receiver = KVCacheStoreRecvingThread( + client, SimpleNamespace(load_mask=lambda hashes, length: [[True]]), + [database], BLOCK, tp_rank=0, ready_event=threading.Event(), + request_level_loads=request_level, recv_workers=workers, + queue_capacity=capacity, + ) + receivers.append(receiver) + return receiver, pool + + yield make + for receiver in receivers: + receiver.stop(cancel_pending=True) + + +@pytest.mark.parametrize( + "outcome", ["miss", "short", "oversized", "incomplete", "native", "no-client", "geometry"], +) +@pytest.mark.parametrize("inline", [False, True]) +def test_all_failures_are_terminal_request_outcomes(make_receiver, outcome, inline): + receiver, _ = make_receiver(None if outcome == "no-client" else MemoryClient(outcome)) + load = request(block_id=0 if outcome == "geometry" else 1) + if inline: + receiver.load_request_sync(load) + else: + receiver.start() + assert receiver.add_request(load) + receiver.request_queue.join() + assert receiver.get_and_clear_block_ids_with_load_errors() == set() + assert receiver.get_and_clear_receive_results() == ({"load"}, {"load"}) + assert receiver.get_and_clear_receive_results() == (set(), set()) + + +def test_full_attention_sync_keeps_block_errors_without_async_completion(make_receiver): + receiver, _ = make_receiver(MemoryClient("miss"), request_level=False) + receiver.load_request_sync(request()) + assert receiver.get_and_clear_block_ids_with_load_errors() == {1} + assert receiver.get_and_clear_receive_results() == (set(), set()) + + +def test_rejection_and_shutdown_cleanup_preserve_paired_outcomes(make_receiver): + receiver, _ = make_receiver(capacity=1) + assert receiver.add_request(request("queued")) + assert not receiver.add_request(request("rejected")) + receiver.stop(cancel_pending=True) + # Shutdown clears ownership states, not the pending failure snapshot. + assert receiver.get_and_clear_receive_results() == ( + {"queued", "rejected"}, {"rejected"}, + ) + assert not receiver.add_request(request("closed")) + assert not receiver.add_request(request("closed")) + assert receiver.get_and_clear_receive_results() == ({"closed"}, {"closed"}) + assert receiver.get_and_clear_receive_results() == (set(), set()) + assert receiver.get_and_clear_block_ids_with_load_errors() == set() + + +def test_native_exception_waits_for_owner_gpu_before_atomic_drain(make_receiver, monkeypatch): + receiver, _ = make_receiver(MemoryClient("native")) + receiver._cuda_device = 7 + fence_entered = threading.Event() + release_fence = threading.Event() + devices = [] + + def synchronize(device): + devices.append(device) + fence_entered.set() + if not release_fence.wait(5): + raise TimeoutError("test did not release the native completion fence") + + monkeypatch.setattr("torch.cuda.synchronize", synchronize) + receiver.start() + try: + receiver.add_request(request()) + assert fence_entered.wait(5) + # Native failure has been recorded, but its GPU ownership is not done. + assert receiver.get_and_clear_receive_results() == (set(), set()) + assert receiver.get_and_clear_block_ids_with_load_errors() == set() + finally: + release_fence.set() + receiver.request_queue.join() + assert devices == [7] + assert receiver.get_and_clear_receive_results() == ({"load"}, {"load"}) + + +@pytest.mark.parametrize("native_failure", [False, True]) +@pytest.mark.parametrize("fail_closed", [False, True]) +@pytest.mark.parametrize("inline", [False, True]) +def test_cancellation_never_releases_active_native_ownership( + make_receiver, native_failure, fail_closed, inline, +): + entered = threading.Event() + release = threading.Event() + cancelled = threading.Event() + + class DelayedClient(MemoryClient): + def batch_get_auto_sg(self, *args): + entered.set() + if not release.wait(5): + raise TimeoutError("test did not release native GET") + return super().batch_get_auto_sg(*args) + + receiver, _ = make_receiver(DelayedClient("native" if native_failure else "ok")) + load_thread = None + if inline: + load_thread = threading.Thread( + target=receiver.load_request_sync, args=(request(),), + ) + load_thread.start() + else: + receiver.start() + receiver.add_request(request()) + assert entered.wait(5) + + def cancel(): + receiver.cancel_requests({"load"}, wait=True, fail_closed=fail_closed) + cancelled.set() + + cancel_thread = threading.Thread(target=cancel) + cancel_thread.start() + try: + assert not cancelled.wait(0.05) + assert receiver.get_and_clear_receive_results() == (set(), set()) + finally: + release.set() + cancel_thread.join(5) + receiver.request_queue.join() + if load_thread is not None: + load_thread.join(5) + assert not load_thread.is_alive() + assert not cancel_thread.is_alive() + assert cancelled.is_set() + assert receiver.get_and_clear_receive_results() == ( + {"load"}, {"load"} if native_failure or fail_closed else set(), + ) + assert receiver.get_and_clear_block_ids_with_load_errors() == set() + + +def test_independent_receives_finish_out_of_order(make_receiver): + slow_entered = threading.Event() + release_slow = threading.Event() + fast_done = threading.Event() + + class OutOfOrderClient(MemoryClient): + def batch_get_auto_sg(self, keys, pointers, capacities): + if not slow_entered.is_set(): + slow_entered.set() + if not release_slow.wait(5): + raise TimeoutError("test did not release the slow GET") + raise RuntimeError("slow receive failed") + result = super().batch_get_auto_sg(keys, pointers, capacities) + fast_done.set() + return result + + receiver, _ = make_receiver(OutOfOrderClient(), workers=2) + receiver.start() + receiver.add_request(request("slow", 1)) + assert slow_entered.wait(5) + receiver.add_request(request("fast", 2)) + try: + assert fast_done.wait(5) + # Join just the fast request's ownership fence, not the shared queue. + receiver.cancel_requests({"fast"}, wait=True, fail_closed=False) + assert receiver.get_and_clear_receive_results() == ({"fast"}, set()) + finally: + release_slow.set() + receiver.request_queue.join() + assert receiver.get_and_clear_receive_results() == ({"slow"}, {"slow"}) + + +def connector_for(receiver, *, inline): + worker = object.__new__(DfkvStoreWorker) + worker.load_async = not inline + worker.request_level_loads = True + worker.kv_recv_thread = receiver + worker.kv_send_thread = None + worker.kv_role = "kv_consumer" + connector = object.__new__(DfkvStoreConnector) + connector.connector_worker = worker + connector._shutdown_condition = threading.Condition() + connector._shutdown = False + connector._inflight_calls = 0 + metadata = DfkvStoreConnectorMetadata({"load"}, set()) + metadata.add_request(request()) + connector.bind_connector_metadata(metadata) + return connector + + +def test_inline_parked_load_fences_both_sides_and_reports_once(make_receiver, monkeypatch): + client = MemoryClient("miss") + receiver, _ = make_receiver(client) + receiver._cuda_device = 7 + calls_seen_at_fence = [] + + def synchronize(device): + assert device == 7 + calls_seen_at_fence.append(client.calls) + + monkeypatch.setattr("torch.cuda.synchronize", synchronize) + connector = connector_for(receiver, inline=True) + connector.start_load_kv(None) + assert client.thread_ids == [threading.get_ident()] + assert calls_seen_at_fence == [0, 1] + result = connector.get_transfer_results(set()) + assert result.finished_recving == result.failed_recving == {"load"} + assert result.finished_sending == set() + # Repeated polling/hooks with the same metadata must not resubmit writes. + connector.start_load_kv(None) + result = connector.get_transfer_results(set()) + assert result.finished_recving == result.failed_recving == set() + assert client.calls == 1 + assert connector._inflight_calls == 0 + assert connector.get_block_ids_with_load_errors() == set() + + +def test_pool_workers_emit_one_completion_and_never_resubmit_polled_metadata(make_receiver): + client = MemoryClient() + receiver, pool = make_receiver(client, workers=3) + connector = connector_for(receiver, inline=False) + receiver.start() + connector.start_load_kv(None) + assert client.calls == 0 + first = connector.get_transfer_results(set()) + receiver.request_queue.join() + second = connector.get_transfer_results(set()) + assert first.finished_recving.isdisjoint(second.finished_recving) + assert first.finished_recving | second.finished_recving == {"load"} + assert first.failed_recving == second.failed_recving == set() + assert ctypes.string_at(ctypes.addressof(pool) + BLOCK, BLOCK) == bytes([17]) * BLOCK + assert client.calls == 1 + assert connector.get_transfer_results(set()).finished_recving == set() + + +def test_connector_failure_update_is_not_skipped_without_kv_events(): + failed = set() + connector = object.__new__(DfkvStoreConnector) + connector.connector_scheduler = SimpleNamespace( + update_connector_output=lambda output: failed.update(output.failed_recving), + ) + connector._kv_cache_events = None + connector.update_connector_output(SimpleNamespace( + failed_recving={"load"}, kv_cache_events=None, + )) + assert failed == {"load"} diff --git a/integration/vllm/tests/test_scheduler_full_block_ids.py b/integration/vllm/tests/test_scheduler_full_block_ids.py index f5d510e..81dbe6a 100644 --- a/integration/vllm/tests/test_scheduler_full_block_ids.py +++ b/integration/vllm/tests/test_scheduler_full_block_ids.py @@ -9,6 +9,7 @@ def test_new_request_metadata_uses_complete_allocated_block_table(): scheduler = object.__new__(DfkvStoreScheduler) + scheduler._failed_load_req_ids = set() scheduler.kv_role = "kv_both" scheduler.client = MagicMock() scheduler.load_specs = {} @@ -49,6 +50,8 @@ def test_new_request_metadata_uses_complete_allocated_block_table(): def test_sampling_tail_does_not_authorize_an_absent_checkpoint(): scheduler = object.__new__(DfkvStoreScheduler) + scheduler._failed_load_req_ids = set() + scheduler.request_level_loads = False scheduler._block_size = 64 scheduler.lookup_async = False scheduler.load_async = False @@ -72,6 +75,8 @@ def test_sampling_tail_does_not_authorize_an_absent_checkpoint(): def test_cache_bypass_discards_stale_external_admission(): scheduler = object.__new__(DfkvStoreScheduler) + scheduler._failed_load_req_ids = set() + scheduler.request_level_loads = False scheduler._block_size = 64 scheduler.lookup_async = False scheduler.load_async = False @@ -97,6 +102,8 @@ def test_consumer_cached_resume_emits_load_without_save( async_pending, same_step_preemption, ): scheduler = object.__new__(DfkvStoreScheduler) + scheduler._failed_load_req_ids = set() + scheduler.request_level_loads = False scheduler.kv_role = "kv_consumer" scheduler.client = MagicMock() scheduler.client.lookup.return_value = 4 @@ -201,6 +208,7 @@ def test_consumer_cached_resume_emits_load_without_save( @pytest.mark.parametrize("async_pending", [False, True]) def test_same_step_preemption_keeps_new_allocation_and_load(async_pending, cached_resume): scheduler = object.__new__(DfkvStoreScheduler) + scheduler._failed_load_req_ids = set() scheduler.kv_role = "kv_both" scheduler.client = MagicMock() scheduler.load_specs = {} @@ -274,3 +282,150 @@ def test_same_step_preemption_keeps_new_allocation_and_load(async_pending, cache step.num_scheduled_tokens = {} metadata = scheduler.build_connector_meta(step) assert metadata.requests == [] + + +@pytest.mark.parametrize( + "group_kinds, load_async, parked", + [ + (["full"], False, False), + (["full"], True, True), + (["full", "full"], False, True), + (["state"], False, True), + (["full", "state"], True, True), + ], +) +def test_load_admission_parks_request_level_layouts( + monkeypatch, group_kinds, load_async, parked, +): + from vllm.v1.kv_cache_interface import FullAttentionSpec + from dfkv_vllm import scheduler as scheduler_module + + client = MagicMock() + monkeypatch.setattr(scheduler_module, "LookupKeyClient", lambda config: client) + monkeypatch.setattr( + scheduler_module, "ensure_deterministic_block_hashing", lambda config: None, + ) + monkeypatch.setattr( + scheduler_module, "resolve_kv_cache_block_sizes", lambda *args: (4, 4), + ) + groups = [ + SimpleNamespace(kv_cache_spec=( + object.__new__(FullAttentionSpec) if kind == "full" else object() + )) + for kind in group_kinds + ] + scheduler = DfkvStoreScheduler( + SimpleNamespace( + cache_config=SimpleNamespace(), + kv_transfer_config=SimpleNamespace( + kv_role="kv_both", kv_connector_extra_config={"load_async": load_async}, + ), + ), + SimpleNamespace(kv_cache_groups=groups, transfer_groups=groups), + ) + request = SimpleNamespace(request_id="load", num_tokens=16, block_hashes=[]) + client.lookup.return_value = 8 + assert scheduler.get_num_new_matched_tokens(request, 0) == (8, parked) + client.lookup.return_value = 0 + assert scheduler.get_num_new_matched_tokens(request, 0) == (0, False) + client.lookup.return_value = None + assert scheduler.get_num_new_matched_tokens(request, 0) == (None, False) + + +@pytest.mark.parametrize("cached_resume", [False, True]) +@pytest.mark.parametrize("load_async", [False, True]) +def test_failed_load_bypasses_persistent_hit_and_rebuilds_computed_save( + cached_resume, load_async, +): + scheduler = object.__new__(DfkvStoreScheduler) + scheduler.kv_role = "kv_both" + scheduler.request_level_loads = True + scheduler.load_async = load_async + scheduler.lookup_async = False + scheduler._block_size = 4 + scheduler.load_specs = {} + scheduler._failed_load_req_ids = set() + scheduler._request_trackers = {} + scheduler._preempted_req_ids = set() + scheduler._unfinished_request_ids = set() + scheduler._unfinished_requests = {} + scheduler._allocated_req_ids = set() + advertised = {"retry": 12} + cached = {} + + def lookup(req_id, *args, **kwargs): + # Discarding the lookup result alone does not repair a remote server + # that still advertises keys whose GET failed. + return cached.setdefault(req_id, advertised[req_id]) + + scheduler.client = SimpleNamespace( + lookup=lookup, discard=lambda req_id: cached.pop(req_id, None), + ) + request = SimpleNamespace( + request_id="retry", num_tokens=20, num_computed_tokens=0, + block_hashes=[], all_token_ids=list(range(20)), + ) + assert scheduler.get_num_new_matched_tokens(request, 0) == (12, True) + old_blocks = ([1, 2, 3], [4, 5, 6]) + scheduler.update_state_after_alloc( + request, SimpleNamespace(get_block_ids=lambda: old_blocks), 12, + ) + step = SimpleNamespace( + finished_req_ids=set(), preempted_req_ids=set(), + scheduled_new_reqs=[], scheduled_cached_reqs=SimpleNamespace(req_ids=[]), + num_scheduled_tokens={}, + ) + pending = scheduler.build_connector_meta(step) + assert pending.requests[0].load_spec.can_load + assert pending.requests[0].can_save is False + + scheduler.update_connector_output(SimpleNamespace(failed_recving={"retry"})) + assert cached == {} + assert scheduler.build_connector_meta(step).requests == [] + assert scheduler.get_num_new_matched_tokens(request, 0) == (0, False) + assert advertised == {"retry": 12} + assert "retry" not in scheduler.load_specs + + # Upstream freed the failed allocation and retries from locally computed + # tokens. MRV1 may resume via cached data without a new preemption event. + new_blocks = ([10], [20]) + scheduler.update_state_after_alloc( + request, SimpleNamespace(get_block_ids=lambda: new_blocks), 0, + ) + step.num_scheduled_tokens = {"retry": 4} + if cached_resume: + step.scheduled_cached_reqs = SimpleNamespace( + req_ids=["retry"], new_block_ids=[new_blocks], num_computed_tokens=[0], + ) + else: + step.scheduled_new_reqs = [SimpleNamespace( + req_id="retry", num_computed_tokens=0, prefill_token_ids=None, + prompt_token_ids=request.all_token_ids, + )] + recomputed = scheduler.build_connector_meta(step).requests[0] + assert recomputed.load_spec is None + assert recomputed.can_save is True + assert recomputed.token_len_chunk == 4 + assert recomputed.block_ids == new_blocks + assert recomputed.token_ids == list(range(4)) + # A later preemption must not re-enable the same failed remote source. + step.preempted_req_ids = {"retry"} + step.scheduled_new_reqs = [] + step.scheduled_cached_reqs = SimpleNamespace(req_ids=[]) + step.num_scheduled_tokens = {} + assert scheduler.build_connector_meta(step).requests == [] + step.preempted_req_ids = set() + assert scheduler.get_num_new_matched_tokens(request, 0) == (0, False) + + + # Failure quarantine ends only at terminal completion. + scheduler.request_finished(request, new_blocks) + step.finished_req_ids = {"retry"} + step.scheduled_new_reqs = [] + step.scheduled_cached_reqs = SimpleNamespace(req_ids=[]) + step.num_scheduled_tokens = {} + assert scheduler.build_connector_meta(step).requests == [] + # Late receive completion of an aborted/finalized request cannot resurrect + # quarantine or stale allocation metadata. + scheduler.update_connector_output(SimpleNamespace(failed_recving={"retry"})) + assert scheduler.get_num_new_matched_tokens(request, 0) == (12, True) diff --git a/integration/vllm/tests/test_stateful_lookup_admission.py b/integration/vllm/tests/test_stateful_lookup_admission.py index c36434d..eb38c5c 100644 --- a/integration/vllm/tests/test_stateful_lookup_admission.py +++ b/integration/vllm/tests/test_stateful_lookup_admission.py @@ -583,7 +583,7 @@ def delayed_put(*args): metadata = SimpleNamespace(requests=[request], preempted_req_ids=set()) def next_model_step(): - producer.get_finished(set(), metadata) + producer.get_transfer_results(set(), metadata) buffers["state"][3].fill_(99) advanced.set() @@ -619,13 +619,13 @@ def test_same_step_resume_save_releases_finished_blocks(hybrid_workers, monkeypa ) metadata = SimpleNamespace(requests=[request], preempted_req_ids={request.req_id}) producer.handle_preemptions(metadata) - producer.get_finished(set(), metadata) + producer.get_transfer_results(set(), metadata) key = PoolKey(producer.token_dbs[1].metadata, hashes[0].hex()).to_bytes() assert objects[producer.client.namespace, key] == bytes([30]) * 16 - done, _ = producer.get_finished( + result = producer.get_transfer_results( {request.req_id}, SimpleNamespace(requests=[], preempted_req_ids=set()), ) - assert done == {request.req_id} + assert result.finished_sending == {request.req_id} def test_cancelled_save_generation_cannot_revive(hybrid_workers, monkeypatch): @@ -680,7 +680,7 @@ def delay_first_exists(keys): new_key = PoolKey(producer.token_dbs[0].metadata, hashes[1].hex()).to_bytes() assert (producer.client.namespace, old_key) not in objects assert objects[producer.client.namespace, new_key] == bytes([22]) * BLOCK - done, _ = producer.get_finished( + result = producer.get_transfer_results( {fresh.req_id}, SimpleNamespace(requests=[], preempted_req_ids=set()), ) - assert done == {fresh.req_id} + assert result.finished_sending == {fresh.req_id} diff --git a/integration/vllm/tests/test_worker_lifecycle.py b/integration/vllm/tests/test_worker_lifecycle.py index ce3716f..c294d63 100644 --- a/integration/vllm/tests/test_worker_lifecycle.py +++ b/integration/vllm/tests/test_worker_lifecycle.py @@ -298,7 +298,7 @@ def test_sync_mode_loads_before_forward(self): self.assertEqual(load.load_spec.token_len, 128) self.assertEqual(skip.load_spec.token_len, 0) - def test_async_mode_defers_load_to_get_finished(self): + def test_async_mode_defers_load_to_result_collection(self): worker = DfkvStoreWorker.__new__(DfkvStoreWorker) worker.load_async = True worker.kv_recv_thread = self.FakeRecv() From fcd5fb4ca125304c0e51b2c8ea85cdf29e28a1b6 Mon Sep 17 00:00:00 2001 From: Ketor Date: Thu, 17 Sep 2026 15:22:50 +0800 Subject: [PATCH 2/3] chore: prepare v2.27.1 hybrid recovery release --- CHANGELOG.md | 18 ++++++++++++++++++ VERSION | 2 +- integration/common/pyproject.toml | 2 +- integration/lmcache/pyproject.toml | 4 ++-- integration/vllm/pyproject.toml | 4 ++-- 5 files changed, 24 insertions(+), 6 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 1e0e1e0..8ba85a5 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,6 +2,24 @@ ## Unreleased +### v2.27.1 — Request-level hybrid KV load recovery + +- Adopt vLLM's native `KVConnectorTransferResults` protocol for hybrid and + multi-group loads instead of reporting ambiguous, group-less block IDs. +- Publish paired receive completion/failure only after native writes and the + owning CUDA device are fenced; cover misses, malformed results, exceptions, + queue rejection and cancellation without losing or duplicating outcomes. +- Keep `load_async=false` I/O serialized on the model thread while parking + hybrid requests until receive completion. Explicit device synchronization + also handles runners that defer the load hook until after a forward pass. +- Bypass failed external hits for the rest of a request and rebuild retry + metadata from fresh allocations, including resumed MRV1 and MRV2 requests. +- Remove the obsolete hybrid invalid-block scheduler patch. This connector + requires native `get_transfer_results` / `failed_recving` support throughout + the vLLM runner, executor and scheduler; older runtimes need an engine upgrade. +- Preserve single full-attention block-level recovery and stored-key layout. + No functional C++/server, native ABI or wire-format changes. + ### v2.27.0 — Runtime RDMA providers and native hybrid-state integration - Include `ibverbs-providers` in the runtime image; exposing RDMA devices does diff --git a/VERSION b/VERSION index 295b40e..c089330 100644 --- a/VERSION +++ b/VERSION @@ -1 +1 @@ -2.27.0 \ No newline at end of file +2.27.1 \ No newline at end of file diff --git a/integration/common/pyproject.toml b/integration/common/pyproject.toml index 0e60f5d..7b23af5 100644 --- a/integration/common/pyproject.toml +++ b/integration/common/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "dfkv-common" -version = "2.27.0" +version = "2.27.1" description = "Canonical namespace and pool-key schema shared by dfkv connectors" requires-python = ">=3.9" diff --git a/integration/lmcache/pyproject.toml b/integration/lmcache/pyproject.toml index e232a47..fff4bd5 100644 --- a/integration/lmcache/pyproject.toml +++ b/integration/lmcache/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "dfkv-connector" -version = "2.27.0" +version = "2.27.1" description = "LMCache RemoteConnector for the dfkv KV cache (ctypes over libdfkv.so)" readme = "README.md" requires-python = ">=3.9" @@ -14,7 +14,7 @@ authors = [{ name = "Wine93", email = "wine93.info@gmail.com" }] # runtime by path (DFKV_LIB / remote_storage_plugin.dfkv.lib). No bundled .so, # no CPython extension, so the wheel is platform-independent. dependencies = [ - "dfkv-common==2.27.0", + "dfkv-common==2.27.1", "lmcache", "torch", ] diff --git a/integration/vllm/pyproject.toml b/integration/vllm/pyproject.toml index 14ee108..08c96a6 100644 --- a/integration/vllm/pyproject.toml +++ b/integration/vllm/pyproject.toml @@ -4,12 +4,12 @@ build-backend = "setuptools.build_meta" [project] name = "dfkv-vllm" -version = "2.27.0" +version = "2.27.1" description = "Direct vLLM KVConnectorBase_V1 connector for dfkv (GPUDirect RDMA, no LMCache)" requires-python = ">=3.12" # vllm + torch are provided by the runtime image; not pinned here. The connector # itself is pure Python (ctypes over libdfkv.so), so there is no native build. -dependencies = ["dfkv-common==2.27.0"] +dependencies = ["dfkv-common==2.27.1"] # Telemetry is opt-in: the OTel SDK is only needed when DFKV_METRICS_ENABLED / # DFKV_TRACE_ENABLED is set. Without this extra the connector stays dependency- From 02bb44d33bf54c39d4d0d087dd5d106a407c8f5a Mon Sep 17 00:00:00 2001 From: Ketor Date: Thu, 17 Sep 2026 15:30:57 +0800 Subject: [PATCH 3/3] fix(vllm): enforce hybrid completion fence despite legacy opt-out --- integration/vllm/README.md | 2 ++ integration/vllm/src/dfkv_vllm/worker.py | 7 +++++-- integration/vllm/tests/test_async_load_device.py | 5 ++++- integration/vllm/tests/test_request_load_failures.py | 1 + 4 files changed, 12 insertions(+), 3 deletions(-) diff --git a/integration/vllm/README.md b/integration/vllm/README.md index 4ec6f85..c478a2c 100644 --- a/integration/vllm/README.md +++ b/integration/vllm/README.md @@ -53,6 +53,8 @@ these layouts. With `kv_load_failure_policy="recompute"`, the engine releases th failed allocation and retries locally; dfkv bypasses the failed remote source for the remainder of that request, including subsequent preemptions. The engine's `"fail"` policy terminates the affected request instead. +The request-level completion fence is mandatory even when the legacy +`DFKV_GPU_LOAD_FENCE=0` override is present. ## Environment variables (engine process) diff --git a/integration/vllm/src/dfkv_vllm/worker.py b/integration/vllm/src/dfkv_vllm/worker.py index 9574a62..80544ea 100644 --- a/integration/vllm/src/dfkv_vllm/worker.py +++ b/integration/vllm/src/dfkv_vllm/worker.py @@ -1451,8 +1451,11 @@ def _handle_request(self, req_meta: ReqMeta): # Fence even a failed/partial native batch before publishing # its terminal outcome and allowing destination block reuse. if ( - os.environ.get("DFKV_GPU_LOAD_FENCE", "1") == "1" - and self._cuda_device is not None + self._cuda_device is not None + and ( + self.request_level_loads + or os.environ.get("DFKV_GPU_LOAD_FENCE", "1") == "1" + ) ): torch.cuda.synchronize(self._cuda_device) diff --git a/integration/vllm/tests/test_async_load_device.py b/integration/vllm/tests/test_async_load_device.py index 42df2d1..de03db6 100644 --- a/integration/vllm/tests/test_async_load_device.py +++ b/integration/vllm/tests/test_async_load_device.py @@ -12,7 +12,9 @@ @pytest.mark.skipif(torch.cuda.device_count() < 2, reason="requires two CUDA devices") -def test_receive_completion_waits_for_the_owning_gpu(): +@pytest.mark.parametrize("request_level_loads", [False, True]) +def test_receive_completion_waits_for_the_owning_gpu(monkeypatch, request_level_loads): + monkeypatch.setenv("DFKV_GPU_LOAD_FENCE", "0" if request_level_loads else "1") for device in (0, 1): with torch.cuda.device(device): torch.cuda._sleep(1) @@ -47,6 +49,7 @@ def batch_get_auto_sg(self, keys, pointers, capacities): receiver = KVCacheStoreRecvingThread( PendingDeviceWrite(), coordinator, [database], 64, tp_rank=0, ready_event=threading.Event(), + request_level_loads=request_level_loads, ) request = ReqMeta( req_id="async-device", token_len_chunk=64, diff --git a/integration/vllm/tests/test_request_load_failures.py b/integration/vllm/tests/test_request_load_failures.py index 7296988..67e4ea0 100644 --- a/integration/vllm/tests/test_request_load_failures.py +++ b/integration/vllm/tests/test_request_load_failures.py @@ -130,6 +130,7 @@ def test_rejection_and_shutdown_cleanup_preserve_paired_outcomes(make_receiver): def test_native_exception_waits_for_owner_gpu_before_atomic_drain(make_receiver, monkeypatch): + monkeypatch.setenv("DFKV_GPU_LOAD_FENCE", "0") receiver, _ = make_receiver(MemoryClient("native")) receiver._cuda_device = 7 fence_entered = threading.Event()