Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
21 commits
Select commit Hold shift + click to select a range
15fff82
Add lookup_scope to skip secondary tier queries per request
ronensc Jul 9, 2026
eed5a2f
Address review: rename primary->cpu
ronensc Jul 13, 2026
819207b
Address review: Implement according to tier filter design
ronensc Jul 16, 2026
dcc1a08
Rename medium->event_medium
ronensc Jul 16, 2026
763098f
Move KV_LOAD_TIERS_KEY to offloading connector
ronensc Jul 16, 2026
f89e326
Address review: reconcile event_medium and filter_medium
ronensc Jul 20, 2026
a741f29
Address review: centralize filter check in TieringOffloadingManager.l…
ronensc Jul 20, 2026
8e899be
Move Locality enum next to Medium enum
ronensc Jul 20, 2026
87f2470
Remove unused consts
ronensc Jul 20, 2026
e9b3634
Revert "Remove unused consts"
ronensc Jul 21, 2026
093c7e9
Address review: Add Medium-to-str translation util in offloading events
ronensc Jul 21, 2026
c714a87
Assert tier.lookup not called when filter excludes tier
ronensc Jul 21, 2026
cd42cd9
Add empty matcher test case
ronensc Jul 21, 2026
922e4af
Address review: Replace Medium.FS and Medium.OBJ with Medium.STORAGE
ronensc Jul 22, 2026
7431155
Address review: Switch dict[str, str] to TierMatcher
ronensc Jul 22, 2026
e1f1a71
Address review: Add _parse_tier_filter
ronensc Jul 22, 2026
5096c33
Address rewview: add locality to SecondaryTierManager base class
ronensc Jul 22, 2026
d990b79
Address review: Replace MEDIUM_FS and MEDIUM_OBJ with MEDIUM_STORAGE
ronensc Jul 22, 2026
e8534f5
Fix test to use Medium.CPU enum instead MEDIUM_CPU string
ronensc Jul 22, 2026
ca1bdb4
Address review: switch logger.warning_one() with logger.waning()
ronensc Jul 22, 2026
67634b8
Address review: treat empty list as deny all tiers
ronensc Jul 23, 2026
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
24 changes: 11 additions & 13 deletions tests/v1/kv_connector/unit/offloading_connector/test_events.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,8 +9,6 @@
from vllm.config import KVEventsConfig, KVTransferConfig
from vllm.distributed.kv_events import (
MEDIUM_CPU,
MEDIUM_FS,
MEDIUM_OBJ,
BlockRemoved,
BlockStored,
)
Expand All @@ -33,14 +31,15 @@
)
from vllm.v1.kv_offload.base import (
Locality,
Medium,
OffloadingEvent,
OffloadingKVEventsConfig,
OffloadKey,
make_offload_key,
)
from vllm.v1.kv_offload.tiering.spec import TieringOffloadingSpec

_CPU_MEDIUM = MEDIUM_CPU
_CPU_MEDIUM = Medium.CPU
_FULL_ATTENTION_EVENT_SPEC = OffloadingEventGroupSpec(
kv_cache_spec_kind=KVCacheSpecKind.FULL_ATTENTION.value,
kv_cache_spec_sliding_window=None,
Expand Down Expand Up @@ -135,7 +134,7 @@ def _record_lookup_chunks(

def _stored_event(
keys: list[OffloadKey],
medium: str = _CPU_MEDIUM,
medium: Medium = _CPU_MEDIUM,
locality: Locality | None = None,
) -> OffloadingEvent:
return OffloadingEvent(
Expand All @@ -148,7 +147,7 @@ def _stored_event(

def _removed_event(
keys: list[OffloadKey],
medium: str = _CPU_MEDIUM,
medium: Medium = _CPU_MEDIUM,
locality: Locality | None = None,
) -> OffloadingEvent:
return OffloadingEvent(
Expand Down Expand Up @@ -181,7 +180,7 @@ def test_take_events_forwards_locality_to_rich_store():

events = list(
tracker.take_events(
[_stored_event([key], locality=Locality.LOCAL, medium=MEDIUM_FS)]
[_stored_event([key], locality=Locality.LOCAL, medium=Medium.STORAGE)]
)
)

Expand All @@ -199,7 +198,7 @@ def test_take_events_forwards_locality_to_placeholder_store():

events = list(
tracker.take_events(
[_stored_event([key], locality=Locality.REMOTE, medium=MEDIUM_FS)]
[_stored_event([key], locality=Locality.REMOTE, medium=Medium.STORAGE)]
)
)

Expand All @@ -216,7 +215,7 @@ def test_take_events_forwards_locality_to_remove():

events = list(
tracker.take_events(
[_removed_event([key], locality=Locality.LOCAL, medium=MEDIUM_FS)]
[_removed_event([key], locality=Locality.LOCAL, medium=Medium.STORAGE)]
)
)

Expand All @@ -240,7 +239,7 @@ def test_take_events_publishes_routable_block_stored():

for i, event in enumerate(batch1):
assert isinstance(event, BlockStored)
assert event.medium == _CPU_MEDIUM
assert event.medium == _CPU_MEDIUM.value
assert event.block_hashes == [_wire_hash(_hash(i))]
assert event.block_size == block_size
assert event.token_ids == list(
Expand Down Expand Up @@ -324,7 +323,7 @@ def test_lookup_promotion_factor_gt_1_store_and_remove():
assert len(removed) == 1
assert isinstance(removed[0], BlockRemoved)
assert removed[0].block_hashes == expected_hashes
assert removed[0].medium == _CPU_MEDIUM
assert removed[0].medium == _CPU_MEDIUM.value
assert removed[0].group_idx == 0
assert not tracker._pending_event_metadata

Expand Down Expand Up @@ -444,12 +443,11 @@ def test_pending_cpu_removal_consumes_hit_backfill_until_next_hit():
]


@pytest.mark.parametrize("medium", [MEDIUM_FS, MEDIUM_OBJ])
def test_secondary_stored_event_does_not_mutate_cpu_metadata(medium: str):
def test_secondary_stored_event_does_not_mutate_cpu_metadata():
tracker, _, _, key = _lookup_chunk()
expected_metadata = dict(tracker._pending_event_metadata)

stored = list(tracker.take_events([_stored_event([key], medium)]))
stored = list(tracker.take_events([_stored_event([key], Medium.STORAGE)]))
assert stored[0].token_ids == [1, 2, 3, 4]
assert tracker._pending_event_metadata == expected_metadata

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@
)
from vllm.v1.kv_offload.base import (
LookupResult,
Medium,
OffloadingEvent,
OffloadingManager,
OffloadPolicy,
Expand Down Expand Up @@ -299,7 +300,7 @@ def take_raw_events():

assert not tracker._pending_event_metadata

raw_events.append(OffloadingEvent(keys=[key], medium=MEDIUM_CPU, removed=False))
raw_events.append(OffloadingEvent(keys=[key], medium=Medium.CPU, removed=False))
events = list(runner.connector_scheduler.take_events())
assert len(events) == 1
assert isinstance(events[0], BlockStored)
Expand All @@ -320,7 +321,7 @@ def take_raw_events():
)
assert key in tracker._pending_event_metadata

raw_events.append(OffloadingEvent(keys=[key], medium=MEDIUM_CPU, removed=True))
raw_events.append(OffloadingEvent(keys=[key], medium=Medium.CPU, removed=True))
[event] = runner.connector_scheduler.take_events()
assert isinstance(event, BlockRemoved)
assert event.medium == MEDIUM_CPU
Expand Down Expand Up @@ -355,7 +356,7 @@ def test_promotion_hit_precedes_stored_event_translation(
raw_events: list[OffloadingEvent] = []

def lookup(key, req_context):
raw_events.append(OffloadingEvent(keys=[key], medium=MEDIUM_CPU, removed=False))
raw_events.append(OffloadingEvent(keys=[key], medium=Medium.CPU, removed=False))
return LookupResult.HIT

def take_raw_events():
Expand Down
4 changes: 2 additions & 2 deletions tests/v1/kv_offload/cpu/test_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,10 +6,10 @@
import numpy as np
import pytest

from vllm.distributed.kv_events import MEDIUM_CPU
from vllm.v1.kv_offload.base import (
LoadStoreSpec,
LookupResult,
Medium,
OffloadingEvent,
OffloadKey,
PrepareStoreOutput,
Expand Down Expand Up @@ -115,7 +115,7 @@ def verify_events(
stores: list[set[OffloadKey]] = []
evictions: list[set[OffloadKey]] = []
for event in events:
assert event.medium == MEDIUM_CPU
assert event.medium == Medium.CPU
if event.removed:
evictions.append(set(event.keys))
else:
Expand Down
7 changes: 3 additions & 4 deletions tests/v1/kv_offload/tiering/test_fs_tier.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,10 +18,10 @@
import pytest
import torch

from vllm.distributed.kv_events import MEDIUM_FS
from vllm.v1.kv_offload.base import (
Locality,
LookupResult,
Medium,
OffloadingEvent,
OffloadingKVEventsConfig,
OffloadKey,
Expand Down Expand Up @@ -514,8 +514,7 @@ def test_successful_store_emits_stored_event(fs_tier_with_events):
events = list(tier.take_events())
assert len(events) == 1
assert events[0].keys == keys
# Literal medium pins the wire contract, not just the constant choice.
assert events[0].medium == "FS"
assert events[0].medium == Medium.STORAGE
assert events[0].locality is Locality.LOCAL
assert not events[0].removed
# take_events drains the buffer.
Expand Down Expand Up @@ -685,7 +684,7 @@ def test_cascade_store_emits_fs_event_through_tiering_manager(tmp_path):
events.extend(manager.take_events())
time.sleep(0.01)

fs_events = [e for e in events if e.medium == MEDIUM_FS]
fs_events = [e for e in events if e.medium == Medium.STORAGE]
assert len(fs_events) == 1
assert set(fs_events[0].keys) == set(keys)
assert not fs_events[0].removed
Expand Down
4 changes: 2 additions & 2 deletions tests/v1/kv_offload/tiering/test_obj_tier.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
from vllm.v1.kv_offload.base import (
Locality,
LookupResult,
Medium,
OffloadingKVEventsConfig,
OffloadKey,
ReqContext,
Expand Down Expand Up @@ -473,8 +474,7 @@ def test_successful_store_emits_stored_event(self):
events = list(self.tier.take_events())
assert len(events) == 1
assert events[0].keys == keys
# Literal medium pins the wire contract, not just the constant choice.
assert events[0].medium == "OBJ"
assert events[0].medium == Medium.STORAGE
assert events[0].locality is Locality.REMOTE
assert not events[0].removed
# take_events drains the buffer.
Expand Down
139 changes: 136 additions & 3 deletions tests/v1/kv_offload/tiering/test_tiering_offloading.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,15 +20,22 @@
from vllm.distributed.kv_transfer.kv_connector.v1.offloading.metrics import (
OffloadingConnectorStats,
)
from vllm.distributed.kv_transfer.kv_connector.v1.offloading.scheduler import (
_parse_tier_filter,
)
from vllm.v1.kv_offload.base import (
Locality,
LookupResult,
Medium,
OffloadingCounterMetadata,
OffloadingEvent,
OffloadKey,
OffloadPolicy,
ReqContext,
RequestOffloadingContext,
ScheduleEndContext,
TierFilter,
TierMatcher,
make_offload_key,
)
from vllm.v1.kv_offload.tiering.base import (
Expand Down Expand Up @@ -245,9 +252,9 @@ def _start_request(self, req_context: ReqContext = _CTX):
self.manager.on_new_request(req_context)

def test_take_events_aggregates_tier_owned_events(self, manager_setup):
primary_event = OffloadingEvent(to_keys([1]), "CPU", removed=False)
secondary_event1 = OffloadingEvent(to_keys([2]), "tier-1", removed=False)
secondary_event2 = OffloadingEvent(to_keys([3]), "tier-2", removed=True)
primary_event = OffloadingEvent(to_keys([1]), Medium.CPU, removed=False)
secondary_event1 = OffloadingEvent(to_keys([2]), Medium.STORAGE, removed=False)
secondary_event2 = OffloadingEvent(to_keys([3]), Medium.STORAGE, removed=True)

self.primary_tier.take_events = MagicMock(return_value=[primary_event])
self.secondary_tier1.take_events = MagicMock(return_value=[secondary_event1])
Expand Down Expand Up @@ -973,6 +980,56 @@ def test_reset_cache_drains_all_tiers(self, manager_setup):
self.secondary_tier2.drain_jobs.assert_called_once()
assert self.manager._transfer_jobs == {}

@pytest.mark.parametrize(
"load_tier_filter",
[
TierFilter(matchers=(TierMatcher(medium=Medium.STORAGE),)),
TierFilter(matchers=()),
],
ids=["non_matching_medium", "empty_no_load"],
)
def test_tier_filter_skips_filtered_secondary(
self, manager_setup, load_tier_filter
):
"""Filter excluding secondary medium returns MISS from secondaries
even when they hold the block; primary is unaffected."""
blocks = to_keys(range(2))
# Put one block in primary, one only in secondary
self._start_request()
self.manager.prepare_store(blocks[:1], _CTX)
self.manager.complete_store(blocks[:1], _CTX, success=True)
self.secondary_tier1.blocks[blocks[1]] = True

# Secondaries have medium=CPU, so load_tier_filter skips them.
self.secondary_tier1.lookup = MagicMock(wraps=self.secondary_tier1.lookup)

ctx = ReqContext(req_id="r1", load_tier_filter=load_tier_filter)
assert self.manager.lookup(blocks[0], ctx) is LookupResult.HIT
assert self.manager.lookup(blocks[1], ctx) is LookupResult.MISS
self.secondary_tier1.lookup.assert_not_called()

@pytest.mark.parametrize(
"load_tier_filter",
[
TierFilter.ALL,
TierFilter(matchers=(TierMatcher(medium=Medium.CPU),)),
TierFilter(matchers=(TierMatcher(),)),
],
ids=["all", "explicit_cpu", "unconstrained_matcher"],
)
def test_tier_filter_allows_matching_secondary(
self, manager_setup, load_tier_filter
):
"""Filter that matches the secondary's medium allows lookup."""
blocks = to_keys(range(1))
self.secondary_tier1.blocks[blocks[0]] = True

self.secondary_tier1.lookup = MagicMock(wraps=self.secondary_tier1.lookup)

ctx = ReqContext(req_id="r2", load_tier_filter=load_tier_filter)
assert self.manager.lookup(blocks[0], ctx) is LookupResult.RETRY
self.secondary_tier1.lookup.assert_called()


class TestTieringOffloadingWithoutSecondaryTiers:
"""Test TieringOffloadingManager with no secondary tiers (backward compat)."""
Expand All @@ -999,5 +1056,81 @@ def test_works_without_secondary_tiers(self):
assert count_hits(manager, blocks) == 3


@pytest.mark.parametrize(
"raw,expected",
[
(
[{"medium": "storage"}],
TierFilter(matchers=(TierMatcher(medium=Medium.STORAGE),)),
),
(
[{"medium": "CPU"}],
TierFilter(matchers=(TierMatcher(medium=Medium.CPU),)),
),
(
[{}],
TierFilter(matchers=(TierMatcher(),)),
),
(
[{"medium": "storage", "locality": "local"}],
TierFilter(
matchers=(TierMatcher(medium=Medium.STORAGE, locality=Locality.LOCAL),)
),
),
(
[{"medium": "cpu"}, {"medium": "storage"}],
TierFilter(
matchers=(
TierMatcher(medium=Medium.CPU),
TierMatcher(medium=Medium.STORAGE),
)
),
),
(
[],
TierFilter(matchers=()),
),
],
ids=[
"medium_storage",
"medium_cpu_uppercase",
"unconstrained",
"with_locality",
"multiple_matchers",
"empty_list_deny_all",
],
)
def test_parse_tier_filter_valid(raw, expected):
assert _parse_tier_filter(raw) == expected


@pytest.mark.parametrize(
"raw",
[
"not a list",
[{"medium": "unknown"}],
[{"locality": "nowhere"}],
],
ids=["non_list", "invalid_medium", "invalid_locality"],
)
def test_parse_tier_filter_invalid_returns_all(raw):
assert _parse_tier_filter(raw) is TierFilter.ALL


def test_parse_tier_filter_skips_bad_entries():
result = _parse_tier_filter(
[
{"medium": "storage"},
"not a dict",
{"medium": "bogus"},
{"medium": "cpu"},
]
)
assert result.matchers == (
TierMatcher(medium=Medium.STORAGE),
TierMatcher(medium=Medium.CPU),
)


if __name__ == "__main__":
pytest.main([__file__, "-v"])
3 changes: 1 addition & 2 deletions vllm/distributed/kv_events.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,8 +44,7 @@ class KVCacheEvent(

MEDIUM_GPU = "GPU"
MEDIUM_CPU = "CPU"
MEDIUM_FS = "FS"
MEDIUM_OBJ = "OBJ"
MEDIUM_STORAGE = "STORAGE"


class BlockStored(KVCacheEvent):
Expand Down
Loading
Loading