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
162 changes: 162 additions & 0 deletions tests/v1/kv_connector/unit/offloading_connector/test_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,11 +14,14 @@
build_offloading_config,
)
from vllm.distributed.kv_transfer.kv_connector.v1.offloading.scheduler import (
OffloadingConnectorScheduler,
RequestOffloadState,
SchedulerOffloadConfig,
is_store_reachable_swa_chunk,
)
from vllm.platforms import current_platform
from vllm.v1.kv_cache_interface import (
CircularBufferSpec,
FullAttentionSpec,
HiddenStateCacheSpec,
KVCacheConfig,
Expand All @@ -31,6 +34,7 @@
SlidingWindowSpec,
UniformTypeKVCacheSpecs,
)
from vllm.v1.kv_offload.base import GPULoadStoreSpec, get_offload_group_idx


def _make_vllm_config(
Expand Down Expand Up @@ -750,3 +754,161 @@ def test_blocks_per_chunk_must_be_positive():

with pytest.raises(ValueError, match="greater than 0"):
build_offloading_config(config, _make_kv_cache_config())


def _make_scratch_hybrid_kv_cache_config() -> KVCacheConfig:
"""Hybrid layout with a non-prefix-cacheable scratch group.

Mirrors sparse-MLA hybrids (e.g. GLM-5.3-Flash) where a CircularBufferSpec
ring holds pre-compression keys: its block size (4) does not divide the
hash granularity of the prefix-cacheable groups (16).
"""
return KVCacheConfig(
num_blocks=4,
kv_cache_tensors=[],
kv_cache_groups=[
KVCacheGroupSpec(["full_layer"], _full_attention_spec()),
KVCacheGroupSpec(
["scratch_layer"],
CircularBufferSpec(
block_size=4,
num_kv_heads=1,
head_size=1,
dtype=torch.float32,
),
),
KVCacheGroupSpec(
["mamba_layer"],
MambaSpec(
block_size=16,
shapes=((1, 1),),
dtypes=(torch.float32,),
mamba_cache_mode="align",
),
),
],
)


def test_scratch_group_does_not_crash_config_translation():
"""A non-prefix-cacheable scratch group whose block size does not divide
tokens_per_hash previously tripped the divisibility assert at boot. It
must be excluded from the offload groups, keeping original indices."""
offloading_config = build_offloading_config(
_make_vllm_config(), _make_scratch_hybrid_kv_cache_config()
)

assert [group.group_idx for group in offloading_config.groups] == [0, 2]
assert tuple(group.tokens_per_block for group in offloading_config.groups) == (
16,
16,
)
assert offloading_config.cache.tokens_per_hash == 16


def test_scratch_group_gets_no_offload_keys():

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 we fold these assertions into test_scratch_group_gets_no_load_slots? It already creates the scheduler, generates keys and checks their group IDs. We can remove this test and the unused RequestOffloadState import, while keeping the config-translation and misalignment tests.

"""The scheduler side must neither look up nor key the scratch group,
while per-group runtime state still spans every KV cache group."""
config = _make_vllm_config()
config.speculative_config = None
kv_cache_config = _make_scratch_hybrid_kv_cache_config()
offloading_config = build_offloading_config(config, kv_cache_config)
spec = MockOffloadingSpec(offloading_config)

scheduler_config = SchedulerOffloadConfig.from_spec(spec, config, kv_cache_config)
assert scheduler_config.num_kv_cache_groups == 3
assert [gc.group_idx for gc in scheduler_config.kv_group_configs] == [0, 2]
assert not scheduler_config.supports_partial_tail

scheduler = OffloadingConnectorScheduler(spec, config, kv_cache_config)
assert scheduler._lookup_groups == (0, 2)

req = MagicMock()
req.kv_transfer_params = None
req.block_hashes = [b"hash-0", b"hash-1"]
req_state = RequestOffloadState(
config=scheduler_config,
req=req,
req_context=MagicMock(),
offloading_context=MagicMock(),
)
assert len(req_state.group_states) == 3

req_state.update_offload_keys()
assert not req_state.group_states[1].offload_keys
assert {
get_offload_group_idx(key)
for group_state in req_state.group_states
for key in group_state.offload_keys
} == {0, 2}


def test_prefix_cacheable_misaligned_group_still_asserts():
"""The divisibility assert must keep guarding prefix-cacheable groups: a
mamba group outside "align" mode backs the hash granularity off to the
scheduler block size (the LCM), which neither group's block size divides."""
kv_cache_config = KVCacheConfig(
num_blocks=4,
kv_cache_tensors=[],
kv_cache_groups=[
KVCacheGroupSpec(["full_layer"], _full_attention_spec()),
KVCacheGroupSpec(
["mamba_layer"],
MambaSpec(
block_size=24,
shapes=((1, 1),),
dtypes=(torch.float32,),
),
),
],
)

with pytest.raises(AssertionError, match="not divisible"):
build_offloading_config(_make_vllm_config(), kv_cache_config)


def test_scratch_group_gets_no_load_slots():
"""update_state_after_alloc must emit full-length GPULoadStoreSpec group
arrays (matching the worker's per-group layout) with a zero-sized entry
for the scratch group, and draw no load keys or destination blocks from
it."""
config = _make_vllm_config()
config.speculative_config = None
kv_cache_config = _make_scratch_hybrid_kv_cache_config()
spec = MockOffloadingSpec(build_offloading_config(config, kv_cache_config))
scheduler = OffloadingConnectorScheduler(spec, config, kv_cache_config)

request = MagicMock()
request.request_id = "req"
request.kv_transfer_params = None
request.block_hashes = [b"hash-0", b"hash-1"]
scheduler.on_new_request(request)
req_status = scheduler._req_status["req"]
req_status.update_offload_keys()
req_status.num_locally_computed_tokens = 0
Comment on lines +875 to +888

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.

The current partial-tail assertion passes even without the new guard because the block and hash sizes are both 16. This combines the key assertions with the load test and uses 4-token hashes to exercise the guard. Tested: 46 config tests pass; removing the guard now fails this test.

Suggested change
config = _make_vllm_config()
config.speculative_config = None
kv_cache_config = _make_scratch_hybrid_kv_cache_config()
spec = MockOffloadingSpec(build_offloading_config(config, kv_cache_config))
scheduler = OffloadingConnectorScheduler(spec, config, kv_cache_config)
request = MagicMock()
request.request_id = "req"
request.kv_transfer_params = None
request.block_hashes = [b"hash-0", b"hash-1"]
scheduler.on_new_request(request)
req_status = scheduler._req_status["req"]
req_status.update_offload_keys()
req_status.num_locally_computed_tokens = 0
config = _make_vllm_config()
config.speculative_config = None
config.cache_config.prefix_match_unit = 4
kv_cache_config = _make_scratch_hybrid_kv_cache_config()
spec = MockOffloadingSpec(build_offloading_config(config, kv_cache_config))
scheduler = OffloadingConnectorScheduler(spec, config, kv_cache_config)
request = MagicMock()
request.request_id = "req"
request.kv_transfer_params = None
request.block_hashes = [f"hash-{i}".encode() for i in range(8)]
scheduler.on_new_request(request)
req_status = scheduler._req_status["req"]
req_status.update_offload_keys()
assert len(req_status.group_states) == 3
assert not req_status.group_states[1].offload_keys
assert scheduler._lookup_groups == (0, 2)
assert not scheduler.config.supports_partial_tail
req_status.num_locally_computed_tokens = 0


def _pending_block(block_id: int) -> MagicMock:
block = MagicMock()
block.block_id = block_id
block.is_null = False
block.block_hash = None
return block

blocks = MagicMock()
blocks.blocks = (
[_pending_block(11), _pending_block(12)], # full attention (group 0)
[_pending_block(31)], # scratch ring (group 1)
[_pending_block(21), _pending_block(22)], # mamba (group 2)
)
scheduler.update_state_after_alloc(request, blocks, num_external_tokens=32)

[load_job] = scheduler._current_batch_load_jobs.values()
dst_spec = load_job.dst_spec
assert isinstance(dst_spec, GPULoadStoreSpec)
assert dst_spec.group_sizes == [2, 0, 2]
assert dst_spec.block_indices == [0, 0, 0]
assert dst_spec.block_ids.tolist() == [11, 12, 21, 22]
assert {get_offload_group_idx(key) for key in load_job.src_spec.offload_keys} == {
0,
2,
}
18 changes: 12 additions & 6 deletions tests/v1/kv_connector/unit/offloading_connector/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -314,12 +314,18 @@ def __init__(
assert isinstance(manager, MagicMock)
self.manager: MagicMock = manager

num_kv_groups = len(kv_cache_config.kv_cache_groups)
assert len(self.connector_scheduler.config.kv_group_configs) == num_kv_groups
for group_config, kv_cache_group in zip(
self.connector_scheduler.config.kv_group_configs,
kv_cache_config.kv_cache_groups,
):
# The connector only builds group configs for prefix-cacheable groups,
# preserving their original group indices.
eligible_groups = [
(group_idx, kv_cache_group)
for group_idx, kv_cache_group in enumerate(kv_cache_config.kv_cache_groups)
if kv_cache_group.kv_cache_spec.prefix_cacheable
]
kv_group_configs = self.connector_scheduler.config.kv_group_configs
assert [group_config.group_idx for group_config in kv_group_configs] == [
group_idx for group_idx, _ in eligible_groups
]
for group_config, (_, kv_cache_group) in zip(kv_group_configs, eligible_groups):
tokens_per_block = kv_cache_group.kv_cache_spec.block_size
assert group_config.tokens_per_block == tokens_per_block
assert group_config.tokens_per_chunk == tokens_per_block * blocks_per_chunk
Expand Down
6 changes: 3 additions & 3 deletions tests/v1/kv_offload/test_factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,7 @@ def _make_offloading_config(
normalized_extra_config["cpu_bytes_to_use"] = cpu_bytes_to_use

if groups is None:
groups = (OffloadingGroupConfig(16, ("layer",)),)
groups = (OffloadingGroupConfig(16, ("layer",), group_idx=0),)

return OffloadingConfig(
groups=groups,
Expand Down Expand Up @@ -493,8 +493,8 @@ def test_offloading_spec_has_replicated_layout_default():

def test_offloading_spec_uses_normalized_chunk_geometry():
groups = (
OffloadingGroupConfig(12, ("full_layer",)),
OffloadingGroupConfig(16, ("mla_layer",)),
OffloadingGroupConfig(12, ("full_layer",), group_idx=0),
OffloadingGroupConfig(16, ("mla_layer",), group_idx=1),
)
spec = _create_spec(
groups=groups,
Expand Down
5 changes: 4 additions & 1 deletion tests/v1/kv_offload/test_file_mapper.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,8 +26,11 @@ def make_mapper_from_offloading_spec(**kwargs) -> FileMapper:
OffloadingGroupConfig(
tokens_per_block=tokens_per_block,
layer_names=(layer_name,),
group_idx=group_idx,
)
for group_idx, (tokens_per_block, layer_name) in enumerate(
kwargs.get("groups", ())
)
for tokens_per_block, layer_name in kwargs.get("groups", ())
),
worker_kv_bytes_per_block=0,
enable_kv_cache_events=False,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,11 @@ def build_offloading_config(
engine_id = kv_transfer_config.engine_id

parallel_config = vllm_config.parallel_config
# Only prefix-cacheable groups can serve offload hits: block hashes are
# computed over prefix-cacheable groups only (resolve_kv_cache_block_sizes),
# so non-cacheable scratch groups (e.g. CircularBufferSpec) have no valid
# hash granularity and must not be offloaded. Original group indices are
# preserved via group_idx.
groups = tuple(
OffloadingGroupConfig(
tokens_per_block=(
Expand All @@ -49,8 +54,10 @@ def build_offloading_config(
)
),
layer_names=tuple(group.layer_names),
group_idx=group_idx,
)
for group in kv_cache_config.kv_cache_groups
for group_idx, group in enumerate(kv_cache_config.kv_cache_groups)
if group.kv_cache_spec.prefix_cacheable
)

_, tokens_per_hash = resolve_kv_cache_block_sizes(kv_cache_config, vllm_config)
Expand Down Expand Up @@ -85,8 +92,8 @@ def build_offloading_config(

assert len(unique_tokens_per_block) == 1, (
"If 'block_size' is specified in kv_connector_extra_config, "
"there must be at least one KV cache group, "
"and all groups must have the same block size."
"there must be at least one prefix-cacheable KV cache group, "
"and all prefix-cacheable groups must have the same block size."
)

tokens_per_block = unique_tokens_per_block.pop()
Expand Down
Loading
Loading