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

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions tests/v1/kv_connector/unit/test_mooncake_store_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -1279,6 +1279,7 @@ def _make_bare_worker(
scheduler_block_size=block_size,
hash_block_size=block_size,
)
worker._init_lookup_key_prefixes()
return worker


Expand Down Expand Up @@ -1346,6 +1347,7 @@ def test_lookup_checks_all_potential_swa_hit_boundaries():
hash_block_size=8,
retention_interval=0,
)
worker._init_lookup_key_prefixes()
# Candidate order: 3 full-attention chunks, then SWA chunks 3, 7, 11.
# Only the first full chunk and the SWA chunk ending at token 32 exist, so
# lookup should recover a 32-token external prefix hit. A sparse
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -132,21 +132,32 @@ def __hash__(self):
)
)

def to_string(self) -> str:
prefix = (
f"{self.key_metadata.cache_prefix}@"
if self.key_metadata.cache_prefix
else ""
)
@staticmethod
def build_prefix(
key_metadata: KeyMetadata,
*,
tp_rank: int | None = None,
pp_rank: int | None = None,
) -> str:
"""Return the stable prefix for a Mooncake pool key."""
prefix = f"{key_metadata.cache_prefix}@" if key_metadata.cache_prefix else ""
return (
f"{prefix}"
f"{self.key_metadata.model_name}"
f"@tp_rank:{self.key_metadata.tp_rank}"
f"@pcp{self.key_metadata.pcp_rank}"
f"@dcp{self.key_metadata.dcp_rank}"
f"@pp_rank:{self.key_metadata.pp_rank}"
f"@group:{self.key_metadata.group_id}"
f"@{self.chunk_hash}"
f"{key_metadata.model_name}"
f"@tp_rank:{key_metadata.tp_rank if tp_rank is None else tp_rank}"
f"@pcp{key_metadata.pcp_rank}"
f"@dcp{key_metadata.dcp_rank}"
f"@pp_rank:{key_metadata.pp_rank if pp_rank is None else pp_rank}"
f"@group:{key_metadata.group_id}"
)

@staticmethod
def build_key_string(key_prefix: str, chunk_hash: str) -> str:
return f"{key_prefix}@{chunk_hash}"

def to_string(self) -> str:
return self.build_key_string(
self.build_prefix(self.key_metadata), self.chunk_hash
)


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1119,6 +1119,20 @@ def __init__(
)
for g_idx, g in enumerate(self._kv_cache_groups)
]
self._init_lookup_key_prefixes()

def _init_lookup_key_prefixes(self) -> None:
"""Precompute per-group key prefixes expanded across TP/PP ranks."""
tp_count = min(self.tp_size, self.num_kv_head)
self._lookup_key_prefixes = tuple(
tuple(
PoolKey.build_prefix(db.metadata, tp_rank=tp, pp_rank=pp)
for tp in range(tp_count)
for pp in range(self.pp_size)
)
for db in self.token_dbs
)
self._lookup_expected_per_key = tp_count * self.pp_size

def register_cross_layers_kv_caches(self, kv_cache: torch.Tensor) -> None:
"""Register a cross-layers KV cache tensor.
Expand Down Expand Up @@ -1381,22 +1395,17 @@ def lookup(self, token_len: int, block_hashes: Sequence[BlockHash]) -> int:
return 0

# Build per-(group, hash) candidate keys expanded across TP/PP.
# candidate_meta[i] is the (group_id, hash_bytes) for candidate_keys[i].
# candidate_meta stores the (group, hash_bytes) for key slice.
candidate_keys: list[str] = []
candidate_meta: list[tuple[int, bytes]] = []
lookup_masks = self.coord.lookup_mask(token_len)
tp_count = min(self.tp_size, self.num_kv_head)
for g_idx, db in enumerate(self.token_dbs):
spec_block_size = db.block_size
lookup_mask = lookup_masks[g_idx]
key_prefixes = self._lookup_key_prefixes[g_idx]
group_hashes = self.coord.block_hashes_for_spec(
block_hashes, self._kv_cache_groups[g_idx].kv_cache_spec
)
metadata_templates = [
dataclasses.replace(db.metadata, tp_rank=tp, pp_rank=pp)
for tp in range(tp_count)
for pp in range(self.pp_size)
]
for chunk_id, h in enumerate(group_hashes):
start_idx = chunk_id * spec_block_size
if start_idx >= token_len:
Expand All @@ -1405,11 +1414,12 @@ def lookup(self, token_len: int, block_hashes: Sequence[BlockHash]) -> int:
chunk_id >= len(lookup_mask) or not lookup_mask[chunk_id]
):
continue
h_hex = h.hex()
h_bytes = bytes(h)
for md in metadata_templates:
candidate_keys.append(PoolKey(md, h_hex).to_string())
candidate_meta.append((g_idx, h_bytes))
hash_hex = h.hex()
for key_prefix in key_prefixes:
candidate_keys.append(
PoolKey.build_key_string(key_prefix, hash_hex)
)
candidate_meta.append((g_idx, bytes(h)))

if not candidate_keys:
return 0
Expand All @@ -1434,12 +1444,15 @@ def lookup(self, token_len: int, block_hashes: Sequence[BlockHash]) -> int:
return 0

# A (group, hash) is "present" only when every TP*PP rank has it.
expected_per_key = max(1, tp_count * self.pp_size)
present_count: dict[tuple[int, bytes], int] = {}
for gh, exists in zip(candidate_meta, res, strict=True):
if exists == 1:
present_count[gh] = present_count.get(gh, 0) + 1
exists_set = {gh for gh, c in present_count.items() if c >= expected_per_key}
ranks_per_candidate = self._lookup_expected_per_key
exists_set = {
(g_idx, hash_bytes)
for i, (g_idx, hash_bytes) in enumerate(candidate_meta)
if all(
res[i * ranks_per_candidate + j] == 1
for j in range(ranks_per_candidate)
)
}

_masks, hit_length = self.coord.find_longest_cache_hit(
block_hashes, token_len, ExternalCachedBlockPool(exists_set)
Expand Down
Loading