Skip to content

Commit 7244ec3

Browse files
committed
fix(kv_pool): cover physical DCP shards in lookup
Signed-off-by: zmc1997 <40617288+zmc1997@users.noreply.github.com>
1 parent ff02be1 commit 7244ec3

3 files changed

Lines changed: 19 additions & 13 deletions

File tree

tests/ut/distributed/ascend_store/test_metadata.py

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -111,12 +111,11 @@ def test_hash_diff(self):
111111
def test_to_string(self):
112112
k = PoolKey(self.meta, "hash1")
113113
s = k.to_string()
114-
self.assertIn("llama", s)
115-
self.assertIn("@pcp2", s)
116-
self.assertIn("@dcp3", s)
117-
self.assertIn("@head_or_tp_rank:1", s)
118-
self.assertIn("@pp_rank:0", s)
119-
self.assertIn("hash1", s)
114+
self.assertEqual(
115+
s,
116+
"llama@pcp:2@dcp:3@head_or_tp_rank:1@pp_rank:0"
117+
"@group:0@cache_role:kv@cache_family:default@hash1",
118+
)
120119

121120
def test_pp_ranks_use_distinct_keys(self):
122121
other_pp_meta = KeyMetadata("llama", 1, 2, 3, 1)
@@ -148,6 +147,7 @@ def test_to_string_contains_layer_id(self):
148147
meta = KeyMetadata("model", 0, 0, 0, 0)
149148
k = LayerPoolKey(meta, "h1", 5)
150149
s = k.to_string()
150+
self.assertIn("@pcp:0@dcp:0", s)
151151
self.assertIn("@layer_id:5", s)
152152
self.assertIn("model", s)
153153
self.assertTrue(s.endswith("@h1"))

vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/metadata.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -115,7 +115,7 @@ def __hash__(self):
115115
def to_string(self):
116116
return (
117117
f"{self.key_metadata.model_name}"
118-
f"@pcp{self.key_metadata.pcp_rank}@dcp{self.key_metadata.dcp_rank}"
118+
f"@pcp:{self.key_metadata.pcp_rank}@dcp:{self.key_metadata.dcp_rank}"
119119
f"@head_or_tp_rank:{self.key_metadata.head_or_tp_rank}"
120120
f"@pp_rank:{self.key_metadata.pp_rank}"
121121
f"@group:{self.key_metadata.kv_cache_group_id}"
@@ -162,7 +162,7 @@ def __hash__(self):
162162
def to_string(self):
163163
return (
164164
f"{self.key_metadata.model_name}"
165-
f"@pcp{self.key_metadata.pcp_rank}@dcp{self.key_metadata.dcp_rank}"
165+
f"@pcp:{self.key_metadata.pcp_rank}@dcp:{self.key_metadata.dcp_rank}"
166166
f"@head_or_tp_rank:{self.key_metadata.head_or_tp_rank}"
167167
f"@group:{self.key_metadata.kv_cache_group_id}"
168168
f"@cache_role:{self.key_metadata.cache_role}"
@@ -327,7 +327,7 @@ def _get_key_prefix(
327327
group_metadata = self.metadata[kv_cache_group_id]
328328
prefix = (
329329
f"{group_metadata.model_name}"
330-
f"@pcp{group_metadata.pcp_rank}@dcp{group_metadata.dcp_rank}"
330+
f"@pcp:{group_metadata.pcp_rank}@dcp:{group_metadata.dcp_rank}"
331331
f"@head_or_tp_rank:{group_metadata.head_or_tp_rank}"
332332
f"@pp_rank:{group_metadata.pp_rank}"
333333
f"@group:{kv_cache_group_id}"

vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/pool_worker.py

Lines changed: 10 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -2211,12 +2211,18 @@ def _replace_key_field(key: str, field: str, value: int) -> str:
22112211
return f"{key[:value_start]}{value}{key[value_end:]}"
22122212

22132213
def _expand_lookup_keys_by_rank(self, keys: list[str], group_id: int) -> list[str]:
2214+
# All-rank KV pool lookup currently assumes PCP=1.
22142215
expanded: list[str] = []
2216+
num_head_or_tp_ranks = self.get_group_tp_size(group_id)
2217+
# Keep each rank shard's block/layer keys contiguous to match
2218+
# lookup_scheduler()'s [rank_shard][block] result slicing.
22152219
for pp_rank in range(self.pp_size):
2216-
for tp_rank in range(self.get_group_tp_size(group_id)):
2217-
for key in keys:
2218-
tp_key = self._replace_key_field(key, "head_or_tp_rank", tp_rank)
2219-
expanded.append(self._replace_key_field(tp_key, "pp_rank", pp_rank))
2220+
for dcp_rank in range(self.dcp_size):
2221+
for head_or_tp_rank in range(num_head_or_tp_ranks):
2222+
for key in keys:
2223+
rank_key = self._replace_key_field(key, "dcp", dcp_rank)
2224+
rank_key = self._replace_key_field(rank_key, "head_or_tp_rank", head_or_tp_rank)
2225+
expanded.append(self._replace_key_field(rank_key, "pp_rank", pp_rank))
22202226
return expanded
22212227

22222228
def _expand_lookup_key_variants(self, key: str, group_id: int, include_all_ranks: bool) -> list[str]:

0 commit comments

Comments
 (0)