Skip to content
Draft
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
158 changes: 158 additions & 0 deletions tests/v1/core/test_kv_cache_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1683,6 +1683,164 @@ def test_resolve_kv_cache_block_sizes_mixed_dcp_replicated_groups():
assert hash_block_size == 64


def test_replicated_mla_uses_lockstep_pool_capacity_and_contiguous_tensors():
vllm_config = SimpleNamespace(
model_config=SimpleNamespace(max_model_len=262144),
parallel_config=SimpleNamespace(
decode_context_parallel_size=4,
prefill_context_parallel_size=1,
),
cache_config=SimpleNamespace(num_gpu_blocks_override=None),
kv_transfer_config=None,
)
specs: dict[str, KVCacheSpec] = {
f"target.{i}": MLAAttentionSpec(
block_size=64,
num_kv_heads=1,
head_size=576,
dtype=torch.uint8,
cache_dtype_str="nvfp4_ds_mla",
)
for i in range(3)
}
specs.update(
{
f"indexer.{i}": MLAAttentionSpec(
block_size=256,
num_kv_heads=1,
head_size=132,
dtype=torch.uint8,
dcp_replicated=True,
)
for i in range(2)
}
)

grouped_specs = kv_cache_utils.group_and_unify_kv_cache_specs(specs, 4, 1)
assert grouped_specs is not None
groups = kv_cache_utils._get_kv_cache_groups_uniform_groups(grouped_specs)
assert all(
isinstance(group.kv_cache_spec, UniformTypeKVCacheSpecs) for group in groups
)
assert kv_cache_utils._use_lockstep_mla_allocation(groups, 4, 1)

bytes_per_pool_block = 3 * (64 * 432) + 2 * (256 * 132)
request_blocks = 262144 // 256
required_memory = bytes_per_pool_block * request_blocks
kv_cache_config = kv_cache_utils.get_kv_cache_config_from_groups(
vllm_config,
groups,
available_memory=required_memory * 2,
)

assert kv_cache_config.num_blocks == request_blocks * 2
assert all(tensor.block_stride == 0 for tensor in kv_cache_config.kv_cache_tensors)
assert sum(tensor.size for tensor in kv_cache_config.kv_cache_tensors) == (
required_memory * 2
)
assert (
kv_cache_utils._max_memory_usage_bytes_from_groups(vllm_config, groups)
== required_memory
)
assert get_max_concurrency_for_kv_cache_config(
vllm_config, kv_cache_config
) == pytest.approx(2.0)


def test_lockstep_mla_predicate_rejects_nonmatching_layouts():
sharded = MLAAttentionSpec(
block_size=64,
num_kv_heads=1,
head_size=128,
dtype=torch.float32,
)
replicated = MLAAttentionSpec(
block_size=256,
num_kv_heads=1,
head_size=128,
dtype=torch.float32,
dcp_replicated=True,
)
groups = [
KVCacheGroupSpec(["target"], sharded),
KVCacheGroupSpec(["indexer"], replicated),
]
assert not kv_cache_utils._use_lockstep_mla_allocation(groups, 1, 1)

mismatched = MLAAttentionSpec(
block_size=64,
num_kv_heads=1,
head_size=128,
dtype=torch.float32,
dcp_replicated=True,
)
mismatched_groups = [groups[0], KVCacheGroupSpec(["indexer"], mismatched)]
assert not kv_cache_utils._use_lockstep_mla_allocation(mismatched_groups, 4, 1)
assert (
kv_cache_utils.group_and_unify_kv_cache_specs(
{"target": sharded, "indexer": mismatched}, 4, 1
)
is None
)

non_mla = FullAttentionSpec(
block_size=256,
num_kv_heads=1,
head_size=128,
dtype=torch.float32,
dcp_replicated=True,
)
assert not kv_cache_utils._use_lockstep_mla_allocation(
[groups[0], KVCacheGroupSpec(["draft"], non_mla)], 4, 1
)


def test_lockstep_mla_equal_page_sizes_use_distinct_tensors():
sharded = MLAAttentionSpec(
block_size=64,
num_kv_heads=1,
head_size=128,
dtype=torch.bfloat16,
)
replicated = MLAAttentionSpec(
block_size=256,
num_kv_heads=1,
head_size=32,
dtype=torch.bfloat16,
dcp_replicated=True,
)
assert sharded.page_size_bytes == replicated.page_size_bytes
groups = [
KVCacheGroupSpec(["target"], sharded),
KVCacheGroupSpec(["indexer"], replicated),
]
vllm_config = SimpleNamespace(
cache_config=SimpleNamespace(num_gpu_blocks_override=None),
parallel_config=SimpleNamespace(
decode_context_parallel_size=4,
prefill_context_parallel_size=1,
),
kv_transfer_config=None,
)
bytes_per_pool_block = sharded.page_size_bytes + replicated.page_size_bytes

kv_cache_config = kv_cache_utils.get_kv_cache_config_from_groups(
vllm_config,
groups,
available_memory=3 * bytes_per_pool_block,
)

assert kv_cache_config.num_blocks == 3
assert [tensor.shared_by for tensor in kv_cache_config.kv_cache_tensors] == [
["target"],
["indexer"],
]
assert [tensor.size for tensor in kv_cache_config.kv_cache_tensors] == [
3 * sharded.page_size_bytes,
3 * replicated.page_size_bytes,
]


def test_dsv4_engine_capacity_uses_worker_kv_cache_config():
from vllm.v1.engine.core import EngineCore

Expand Down
177 changes: 177 additions & 0 deletions tests/v1/core/test_prefix_caching.py
Original file line number Diff line number Diff line change
Expand Up @@ -177,6 +177,183 @@ def make_kv_cache_config_hybrid_model(
)


def make_lockstep_mla_manager(num_blocks: int = 5) -> KVCacheManager:
global_block_size = 256
kv_cache_config = KVCacheConfig(
num_blocks=num_blocks,
kv_cache_tensors=[],
kv_cache_groups=[
KVCacheGroupSpec(
["target"],
MLAAttentionSpec(
block_size=global_block_size // 4,
num_kv_heads=1,
head_size=432,
dtype=torch.uint8,
),
),
KVCacheGroupSpec(
["indexer"],
MLAAttentionSpec(
block_size=global_block_size,
num_kv_heads=1,
head_size=132,
dtype=torch.uint8,
dcp_replicated=True,
),
),
],
)
return KVCacheManager(
kv_cache_config=kv_cache_config,
max_model_len=4 * global_block_size,
scheduler_block_size=global_block_size,
hash_block_size=global_block_size,
enable_caching=True,
dcp_world_size=4,
)


def test_mixed_mla_groups_share_block_ids_hashes_and_eviction_order():
block_size = 256
manager = make_lockstep_mla_manager()
assert manager.coordinator.lockstep_mla_allocations
request = make_request(
"lockstep",
list(range(4 * block_size)),
block_size,
sha256,
)

new_blocks = manager.allocate_slots(
request,
num_new_tokens=4 * block_size,
full_sequence_must_fit=True,
)
assert new_blocks is not None
target_ids, indexer_ids = manager.get_block_ids(request.request_id)
assert target_ids == indexer_ids
assert new_blocks.get_block_ids() == (target_ids, target_ids)
assert manager.block_pool.get_num_free_blocks() == 0
assert all(
manager.block_pool.blocks[block_id].ref_cnt == 2 for block_id in target_ids
)

for block_hash, block_id in zip(request.block_hashes, target_ids):
group_keys = [
make_block_hash_with_group_id(block_hash, group_id) for group_id in range(2)
]
assert group_keys[0] != group_keys[1]
cached = manager.block_pool.get_cached_block(block_hash, [0, 1])
assert cached is not None
assert [block.block_id for block in cached] == [block_id, block_id]

manager.free(request)
assert manager.block_pool.get_num_free_blocks() == 4
assert all(
manager.block_pool.blocks[block_id].ref_cnt == 0 for block_id in target_ids
)
free_ids = [
block.block_id
for block in manager.block_pool.free_block_queue.get_all_free_blocks()
]
assert free_ids == target_ids[::-1]

replay = make_request(
"lockstep-replay",
list(range(4 * block_size)),
block_size,
sha256,
)
computed_blocks, num_computed_tokens, _ = manager.get_computed_blocks(replay)
assert num_computed_tokens == 3 * block_size
assert computed_blocks.get_block_ids() == (
target_ids[:3],
target_ids[:3],
)

replay_new_blocks = manager.allocate_slots(
replay,
num_new_tokens=block_size,
num_new_computed_tokens=num_computed_tokens,
new_computed_blocks=computed_blocks,
)
assert replay_new_blocks is not None
assert manager.get_block_ids(replay.request_id) == (target_ids, target_ids)
assert all(
manager.block_pool.blocks[block_id].ref_cnt == 2 for block_id in target_ids
)
manager.free(replay)
assert [
block.block_id
for block in manager.block_pool.free_block_queue.get_all_free_blocks()
] == target_ids[::-1]

evicted = manager.block_pool.get_new_blocks(1)[0]
assert evicted.block_id == target_ids[-1]
assert manager.block_pool.get_cached_block(request.block_hashes[-1], [0]) is None
assert manager.block_pool.get_cached_block(request.block_hashes[-1], [1]) is None


def test_lockstep_mla_rejects_external_computed_blocks():
manager = make_lockstep_mla_manager()

with pytest.raises(NotImplementedError, match="External KV loads"):
manager.coordinator.allocate_new_computed_blocks(
"external",
([], []),
num_local_computed_tokens=0,
num_external_computed_tokens=256,
)

assert manager.get_block_ids("external") == ([], [])


def test_lockstep_group_hashes_promote_partial_block_together():
hash_block_size = 64
block_size = 4 * hash_block_size
pool = BlockPool(
num_gpu_blocks=2,
enable_caching=True,
hash_block_size=hash_block_size,
)
block = pool.get_new_blocks(1)[0]
request = make_request(
"promotion",
list(range(block_size)),
hash_block_size,
sha256,
)

for group_id in range(2):
pool.cache_partial_block(
request=request,
block=block,
num_tokens=2 * hash_block_size,
kv_cache_group_id=group_id,
block_size=block_size,
)
partial_hash = request.block_hashes[1]
partial_cached = pool.get_cached_block(partial_hash, [0, 1])
assert partial_cached == [block, block]

for group_id in range(2):
pool.cache_full_blocks(
request=request,
blocks=[block],
num_cached_blocks=0,
num_full_blocks=1,
block_size=block_size,
kv_cache_group_id=group_id,
)

assert pool.get_cached_block(partial_hash, [0]) is None
assert pool.get_cached_block(partial_hash, [1]) is None
full_hash = request.block_hashes[-1]
assert pool.get_cached_block(full_hash, [0, 1]) == [block, block]
assert block.block_hash_num_tokens == block_size


def make_kv_cache_config_three_types(
block_size: int, num_blocks: int, third_spec_type: str = "mamba"
) -> KVCacheConfig:
Expand Down
Loading
Loading