diff --git a/tests/v1/kv_connector/unit/test_invalid_blocks_correctness.py b/tests/v1/kv_connector/unit/test_invalid_blocks_correctness.py index 77d629729776..9154ff564f0c 100644 --- a/tests/v1/kv_connector/unit/test_invalid_blocks_correctness.py +++ b/tests/v1/kv_connector/unit/test_invalid_blocks_correctness.py @@ -56,6 +56,130 @@ def recompute_scheduler(): return create_scheduler(vllm_config) +def _make_recovery_scheduler( + block_ids_by_group: tuple[list[int], ...], + group_block_sizes: tuple[int, ...], + get_num_skipped_tokens: tuple[Callable[[int], int], ...], + scheduler_block_size: int = 16, +) -> Scheduler: + scheduler = object.__new__(Scheduler) + scheduler.kv_cache_manager = Mock() + scheduler.kv_cache_manager.get_block_ids.return_value = block_ids_by_group + + coordinator = Mock() + coordinator.scheduler_block_size = scheduler_block_size + managers = [] + for block_size, get_skipped in zip( + group_block_sizes, get_num_skipped_tokens, strict=True + ): + manager = Mock() + manager.block_size = block_size + manager.get_num_skipped_tokens.side_effect = get_skipped + managers.append(manager) + coordinator.single_type_managers = tuple(managers) + scheduler.kv_cache_manager.coordinator = coordinator + return scheduler + + +def _make_recovery_request(num_computed_tokens: int) -> Mock: + request = Mock() + request.request_id = "req-0" + request.num_computed_tokens = num_computed_tokens + return request + + +def test_hybrid_recovery_uses_actual_failure_and_common_alignment(): + scheduler = _make_recovery_scheduler( + block_ids_by_group=([1, 2], [10, 11, 12, 13]), + group_block_sizes=(16, 8), + get_num_skipped_tokens=(lambda _: 0, lambda _: 0), + ) + request = _make_recovery_request(32) + + affected, tokens, blocks = scheduler._update_requests_with_invalid_blocks( + [request], {13}, {}, evict_blocks=True + ) + + assert affected == {"req-0"} + assert request.num_computed_tokens == 16 + assert tokens == 16 + assert blocks == {2, 12, 13} + + +def test_hybrid_recovery_rechecks_shifted_sliding_window(): + window_size = 32 + scheduler = _make_recovery_scheduler( + block_ids_by_group=([1, 2, 3, 4], [0, 0, 10, 11]), + group_block_sizes=(16, 16), + get_num_skipped_tokens=( + lambda _: 0, + lambda num_tokens: max(0, num_tokens - window_size + 1), + ), + ) + request = _make_recovery_request(64) + + affected, tokens, _ = scheduler._update_requests_with_invalid_blocks( + [request], {11}, {}, evict_blocks=False + ) + + assert affected == {"req-0"} + assert request.num_computed_tokens == 0 + assert tokens == 64 + + +def test_hybrid_recovery_finds_previous_mamba_state(): + scheduler = _make_recovery_scheduler( + block_ids_by_group=([1, 2, 3, 4], [0, 0, 20, 21]), + group_block_sizes=(16, 16), + get_num_skipped_tokens=(lambda _: 0, lambda num_tokens: num_tokens - 1), + ) + request = _make_recovery_request(64) + + affected, tokens, _ = scheduler._update_requests_with_invalid_blocks( + [request], {21}, {}, evict_blocks=False + ) + + assert affected == {"req-0"} + assert request.num_computed_tokens == 48 + assert tokens == 16 + + +def test_mamba_recovery_uses_scheduler_block_alignment(): + scheduler = _make_recovery_scheduler( + block_ids_by_group=([1, 2, 3, 4], [0, 20]), + group_block_sizes=(16, 32), + get_num_skipped_tokens=(lambda _: 0, lambda num_tokens: num_tokens - 1), + scheduler_block_size=32, + ) + request = _make_recovery_request(64) + + affected, tokens, _ = scheduler._update_requests_with_invalid_blocks( + [request], {4}, {}, evict_blocks=False + ) + + assert affected == {"req-0"} + assert request.num_computed_tokens == 0 + assert tokens == 64 + + +def test_load_failure_recovery_does_not_apply_eagle_drop_again(): + scheduler = _make_recovery_scheduler( + block_ids_by_group=([1, 2, 3],), + group_block_sizes=(16,), + get_num_skipped_tokens=(lambda _: 0,), + ) + scheduler.kv_cache_manager.coordinator.single_type_managers[0].use_eagle = True + request = _make_recovery_request(48) + + affected, tokens, _ = scheduler._update_requests_with_invalid_blocks( + [request], {2}, {}, evict_blocks=False + ) + + assert affected == {"req-0"} + assert request.num_computed_tokens == 16 + assert tokens == 32 + + def test_sync_recompute_blocks_not_freed_for_running_requests( recompute_scheduler: Scheduler, ): diff --git a/vllm/v1/core/sched/scheduler.py b/vllm/v1/core/sched/scheduler.py index c08a1302131e..199cfbd0580c 100644 --- a/vllm/v1/core/sched/scheduler.py +++ b/vllm/v1/core/sched/scheduler.py @@ -2751,6 +2751,51 @@ def _update_from_kv_xfer_finished(self, kv_connector_output: KVConnectorOutput): assert req_id in self.requests self._free_blocks(self.requests[req_id]) + def _find_longest_valid_kv_prefix( + self, + block_ids_by_group: tuple[list[int], ...], + invalid_block_ids: set[int], + max_num_computed_tokens: int, + ) -> int: + """Find the longest group-consistent prefix after a KV load failure.""" + coordinator = self.kv_cache_manager.coordinator + managers = coordinator.single_type_managers + assert len(block_ids_by_group) == len(managers) + + alignment = coordinator.scheduler_block_size + candidate = max_num_computed_tokens // alignment * alignment + + missing_prefixes: list[list[int]] = [] + for block_ids in block_ids_by_group: + missing_prefix = [0] + for block_id in block_ids: + is_missing = block_id <= 0 or block_id in invalid_block_ids + missing_prefix.append(missing_prefix[-1] + is_missing) + missing_prefixes.append(missing_prefix) + + while candidate > 0: + for manager, block_ids, missing_prefix in zip( + managers, + block_ids_by_group, + missing_prefixes, + strict=True, + ): + block_size = manager.block_size + first_required = max( + 0, + manager.get_num_skipped_tokens(candidate) // block_size, + ) + end_required = (candidate + block_size - 1) // block_size + if end_required > len(block_ids) or ( + missing_prefix[end_required] != missing_prefix[first_required] + ): + candidate -= alignment + break + else: + return candidate + + return 0 + def _update_requests_with_invalid_blocks( self, requests: Iterable[Request], @@ -2790,67 +2835,67 @@ def _update_requests_with_invalid_blocks( # it. This set tracks blocks already marked for recomputation. marked_invalid_block_ids: set[int] = set() for request in requests: - is_affected = False - marked_invalid_block = False req_id = request.request_id - # TODO (davidb): add support for hybrid memory allocator - (req_block_ids,) = self.kv_cache_manager.get_block_ids(req_id) + req_block_ids_by_group = self.kv_cache_manager.get_block_ids(req_id) # We iterate only over blocks that may contain externally computed # tokens req_num_computed_tokens = ( request.num_computed_tokens - num_scheduled_tokens.get(req_id, 0) ) - req_num_computed_blocks = ( - req_num_computed_tokens + self.block_size - 1 - ) // self.block_size - for idx, block_id in zip(range(req_num_computed_blocks), req_block_ids): - if block_id not in invalid_block_ids: - continue - - is_affected = True - - if block_id in marked_invalid_block_ids: - # This invalid block is shared with a previous request - # and was already marked for recomputation. - # This means this request can still consider this block - # as computed when rescheduled. - # Currently this only applies to sync loading; Async - # loading does not yet support block sharing - continue + request_invalid_block_ids: set[int] = set() + managers = self.kv_cache_manager.coordinator.single_type_managers + for manager, block_ids in zip( + managers, req_block_ids_by_group, strict=True + ): + num_computed_blocks = ( + req_num_computed_tokens + manager.block_size - 1 + ) // manager.block_size + request_invalid_block_ids.update( + block_id + for block_id in block_ids[:num_computed_blocks] + if block_id in invalid_block_ids + ) - marked_invalid_block_ids.add(block_id) + if not request_invalid_block_ids: + continue - if marked_invalid_block: - # This request has already marked an invalid block for - # recomputation and updated its num_computed_tokens. - continue + owned_invalid_block_ids = ( + request_invalid_block_ids - marked_invalid_block_ids + ) + marked_invalid_block_ids.update(request_invalid_block_ids) - marked_invalid_block = True - # Truncate the computed tokens at the first failed block - request.num_computed_tokens = idx * self.block_size - num_affected_tokens = ( - req_num_computed_tokens - request.num_computed_tokens + if not owned_invalid_block_ids: + # All invalid blocks of this request are shared with previous + # requests and will be recomputed by them. + total_affected_tokens += ( + request.num_computed_tokens - req_num_computed_tokens + ) + request.num_computed_tokens = req_num_computed_tokens + else: + rewind_num_computed_tokens = self._find_longest_valid_kv_prefix( + req_block_ids_by_group, + owned_invalid_block_ids, + req_num_computed_tokens, + ) + request.num_computed_tokens = rewind_num_computed_tokens + total_affected_tokens += ( + req_num_computed_tokens - rewind_num_computed_tokens ) - total_affected_tokens += num_affected_tokens - # collect invalid block and all downstream dependent blocks if evict_blocks: - blocks_to_evict.update(req_block_ids[idx:]) - - if is_affected: - if not marked_invalid_block: - # All invalid blocks of this request are shared with - # previous requests and will be recomputed by them. - # Revert to considering only cached tokens as computed. - # Currently this only applies to sync loading; Async - # loading does not yet support block sharing - total_affected_tokens += ( - request.num_computed_tokens - req_num_computed_tokens - ) - request.num_computed_tokens = req_num_computed_tokens + blocks_to_evict.update(request_invalid_block_ids) + for manager, block_ids in zip( + managers, req_block_ids_by_group, strict=True + ): + first_block = rewind_num_computed_tokens // manager.block_size + blocks_to_evict.update( + block_id + for block_id in block_ids[first_block:] + if block_id > 0 + ) - affected_req_ids.add(request.request_id) + affected_req_ids.add(request.request_id) return affected_req_ids, total_affected_tokens, blocks_to_evict