Skip to content
Open
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
109 changes: 109 additions & 0 deletions tests/v1/kv_connector/unit/test_remote_decode_lifecycle.py
Original file line number Diff line number Diff line change
Expand Up @@ -258,3 +258,112 @@ def test_abort_during_kv_transfer():
)
scheduler.update_from_output(scheduler_output, model_runner_output)
assert_scheduler_empty(scheduler)


@pytest.mark.parametrize(
"num_tokens,token_budget,num_lookahead_tokens",
[
(40, 24, 0), # chunk 2 = tokens 24..39: rest of b1 plus b2
(20, 18, 0), # chunk 2 = tokens 18..19: fits in b1, no new block
(40, 30, 3), # spec decode: chunk 1 also allocates lookahead block b2
],
)
def test_host_buffer_save_resaves_block_straddling_chunk_boundary(
num_tokens: int, token_budget: int, num_lookahead_tokens: int

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Could you add unit tests for (1) preempt then resume case and (2) a hybrid (attention + mamba) case?

):
"""Host-buffer mode (kv_buffer_device="cpu") copies each prefill step's
blocks to host memory. When a chunk boundary is not block aligned, the
next chunk writes the rest of the previous chunk's last block, so that
block must be copied again. Blocks allocated ahead of the written tokens
(spec-decode lookahead) must not be copied before they are written."""
block_size = 16
vllm_config = create_vllm_config(
block_size=block_size,
max_num_batched_tokens=token_budget,
kv_role="kv_producer",
)
scheduler = create_scheduler(vllm_config)
scheduler.num_lookahead_tokens = num_lookahead_tokens
connector_scheduler = scheduler.get_kv_connector().connector_scheduler
connector_scheduler.use_host_buffer = True
request = create_request(
request_id=1,
num_tokens=num_tokens,
block_size=block_size,
do_remote_decode=True,
)
req_id = request.request_id
scheduler.add_request(request)

out1 = scheduler.schedule()
(saved1,) = out1.kv_connector_metadata.reqs_to_save[req_id].local_block_ids
(table,) = scheduler.kv_cache_manager.get_block_ids(req_id)
# Only the blocks holding tokens written in chunk 1.
assert saved1 == table[: -(-token_budget // block_size)]
model_output = create_model_runner_output([request])
model_output.sampled_token_ids = [[]]
scheduler.update_from_output(out1, model_output)
assert request.num_computed_tokens % block_size != 0
first = request.num_computed_tokens // block_size
straddling = table[first]

out2 = scheduler.schedule()
(saved2,) = out2.kv_connector_metadata.reqs_to_save[req_id].local_block_ids
(table,) = scheduler.kv_cache_manager.get_block_ids(req_id)
assert saved2[0] == straddling
assert saved2 == table[first : -(-num_tokens // block_size)]
assert req_id not in connector_scheduler._reqs_need_save
assert req_id not in connector_scheduler._reqs_save_state


def test_host_buffer_save_includes_prefix_cache_hit_blocks():
"""A first chunk that starts after a local prefix-cache hit still copies
the cached blocks: D pulls the whole prompt from the host buffer."""
block_size = 16
vllm_config = create_vllm_config(
block_size=block_size,
max_num_batched_tokens=64,
kv_role="kv_producer",
)
scheduler = create_scheduler(vllm_config)
connector_scheduler = scheduler.get_kv_connector().connector_scheduler
connector_scheduler.use_host_buffer = True

# Request 1 computes and caches a 32-token (2-block) prefix.
first = create_request(
request_id=1,
num_tokens=40,
common_prefix_len=32,
block_size=block_size,
do_remote_decode=True,
)
scheduler.add_request(first)
out = scheduler.schedule()
(cached_prefix,) = scheduler.kv_cache_manager.get_block_ids(first.request_id)
cached_prefix = cached_prefix[:2]
scheduler.update_from_output(out, create_model_runner_output([first]))

# Request 2 hits the prefix and needs two chunks for the rest.
second = create_request(
request_id=2,
num_tokens=110,
common_prefix_len=32,
block_size=block_size,
do_remote_decode=True,
)
req_id = second.request_id
scheduler.add_request(second)
out1 = scheduler.schedule()
(saved1,) = out1.kv_connector_metadata.reqs_to_save[req_id].local_block_ids
(table,) = scheduler.kv_cache_manager.get_block_ids(req_id)
assert table[:2] == cached_prefix # the prefix came from the cache
assert saved1 == table[: (32 + 64) // block_size]
model_output = create_model_runner_output([second])
model_output.sampled_token_ids = [[]]
scheduler.update_from_output(out1, model_output)

out2 = scheduler.schedule()
(saved2,) = out2.kv_connector_metadata.reqs_to_save[req_id].local_block_ids
(table,) = scheduler.kv_cache_manager.get_block_ids(req_id)
assert saved2 == table[96 // block_size : -(-110 // block_size)]
assert req_id not in connector_scheduler._reqs_save_state
2 changes: 2 additions & 0 deletions tests/v1/kv_connector/unit/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -558,6 +558,7 @@ def make_nixl_scheduler(
sched._reqs_in_batch = set()
sched._reqs_not_processed = set()
sched._reqs_need_save = {}
sched._reqs_save_state = {}
sched.use_host_buffer = False
sched.engine_id = "test-engine"
sched.transfer_tp_size = 1
Expand Down Expand Up @@ -596,6 +597,7 @@ def make_nixl_push_scheduler(
sched._reqs_in_batch = set()
sched._reqs_not_processed = set()
sched._reqs_need_save = {}
sched._reqs_save_state = {}
sched._kv_lease_duration = 30
sched.decoder_kv_blocks_ttl = decoder_kv_blocks_ttl
sched.use_host_buffer = False
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -117,6 +117,9 @@ def __init__(
ReqId, tuple[Request, BlockIds, tuple[int, ...], bool]
] = {}
self._reqs_need_save: dict[ReqId, Request] = {}
# Host-buffer save progress of a partially prefilled request: its block
# table per KV cache group and how many tokens have been saved.
self._reqs_save_state: dict[ReqId, tuple[list[list[int]], int]] = {}
# Reqs to send and their expiration time
self._reqs_need_send: dict[ReqId, float] = {}
self._reqs_in_batch: set[ReqId] = set()
Expand Down Expand Up @@ -154,6 +157,16 @@ def __init__(
else None
for g in kv_cache_config.transfer_groups
]
# Tokens per block for groups whose blocks map to token positions (one
# entry per KV cache group, before transfer-group selection). None for
# other groups (e.g. SSM state), whose host-buffer save stays per-step.
dcp_size = parallel_config.decode_context_parallel_size
self._save_block_size: list[int | None] = [
spec.block_size * (dcp_size if spec.dcp_sharded else 1)
if isinstance(spec, (FullAttentionSpec, SlidingWindowSpec))
else None
for spec in (g.kv_cache_spec for g in kv_cache_config.kv_cache_groups)
]

# Threshold to decide whether to compute kv cache locally
# or pull from a remote node: minimum number of remote
Expand Down Expand Up @@ -433,34 +446,55 @@ def _build_save_meta(
# only called when use_host_buffer is True to build the save metadata

# NOTE: For the prefill side, there might be a chance that an early added
# request is a chunked prefill, so we need to check if new blocks are added
for req_id, new_block_id_groups, _ in yield_req_data(scheduler_output):
req_to_save = self._reqs_need_save.get(req_id)
if req_to_save is None or new_block_id_groups is None:
# request is a chunked prefill, so we need to check if new blocks are added.
# Blocks are selected by token position: a chunk boundary that is not
# block aligned leaves a block partly written for the next chunk to
# finish, and blocks allocated ahead (e.g. spec-decode lookahead) are
# not written yet.
assert scheduler_output.num_scheduled_tokens is not None
for req_id, new_block_id_groups, resumed in yield_req_data(scheduler_output):
req = self._reqs_need_save.get(req_id)
if req is None:
continue
req = req_to_save
# New and resumed requests send their full block table.
state = None if resumed else self._reqs_save_state.get(req_id)
if state is None:
if new_block_id_groups is None:
continue
table = [list(blocks) for blocks in new_block_id_groups]
saved = 0
else:
table, saved = state
if new_block_id_groups is not None:
for blocks, new_blocks in zip(table, new_block_id_groups):
blocks.extend(new_blocks)
num_scheduled_tokens = scheduler_output.num_scheduled_tokens[req_id]
end = req.num_computed_tokens + num_scheduled_tokens
to_save: list[list[int]] = []
for group_id, blocks in enumerate(table):
block_size = self._save_block_size[group_id]
if block_size is not None:
to_save.append(blocks[saved // block_size : cdiv(end, block_size)])
elif new_block_id_groups is not None:
to_save.append(list(new_block_id_groups[group_id]))
else:
to_save.append([])

assert req.kv_transfer_params is not None
clipped_block_id_groups = self.get_exchange_clipped_blocks(
new_block_id_groups, clip_ssm=False
)
meta.add_new_req_to_save(
request_id=req_id,
local_block_ids=clipped_block_id_groups,
local_block_ids=self.get_exchange_clipped_blocks(
tuple(to_save), clip_ssm=False
),
kv_transfer_params=req.kv_transfer_params,
)
assert scheduler_output.num_scheduled_tokens is not None
num_scheduled_tokens = scheduler_output.num_scheduled_tokens[req_id]
is_partial = (
req.num_computed_tokens + num_scheduled_tokens
) < req.num_prompt_tokens
if not is_partial:
# For non-partial prefills, once new req_meta is scheduled, it
# can be removed from _reqs_need_save.
# For partial prefill case, we will retain the request in
# _reqs_need_save until all blocks are scheduled with req_meta.
# Therefore, only pop if `not is_partial`.
if end < req.num_prompt_tokens:
# Partial prefill: keep the request in _reqs_need_save until
# its last chunk is scheduled.
self._reqs_save_state[req_id] = (table, end)
else:
self._reqs_need_save.pop(req_id)
self._reqs_save_state.pop(req_id, None)

def build_connector_meta(
self,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -241,6 +241,7 @@ def request_finished(
self._reqs_not_processed.add(request.request_id)
# Clear _reqs_need_save if a request is aborted as partial prefill.
self._reqs_need_save.pop(request.request_id, None)
self._reqs_save_state.pop(request.request_id, None)
return False, None

# TODO: check whether block_ids actually ever be 0. If not we could
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -255,6 +255,7 @@ def request_finished(
):
self._reqs_not_processed.add(request.request_id)
self._reqs_need_save.pop(request.request_id, None)
self._reqs_save_state.pop(request.request_id, None)
return False, None

delay_free_blocks = any(len(group) > 0 for group in block_ids)
Expand Down
Loading