Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
124 changes: 124 additions & 0 deletions tests/v1/kv_connector/unit/test_invalid_blocks_correctness.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
):
Expand Down
139 changes: 92 additions & 47 deletions vllm/v1/core/sched/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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],
Expand Down Expand Up @@ -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

Expand Down
Loading