There was an error while loading. Please reload this page.
1 parent ff02be1 commit 7244ec3Copy full SHA for 7244ec3
3 files changed
tests/ut/distributed/ascend_store/test_metadata.py
@@ -111,12 +111,11 @@ def test_hash_diff(self):
111
def test_to_string(self):
112
k = PoolKey(self.meta, "hash1")
113
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)
+ self.assertEqual(
+ s,
+ "llama@pcp:2@dcp:3@head_or_tp_rank:1@pp_rank:0"
+ "@group:0@cache_role:kv@cache_family:default@hash1",
+ )
120
121
def test_pp_ranks_use_distinct_keys(self):
122
other_pp_meta = KeyMetadata("llama", 1, 2, 3, 1)
@@ -148,6 +147,7 @@ def test_to_string_contains_layer_id(self):
148
147
meta = KeyMetadata("model", 0, 0, 0, 0)
149
k = LayerPoolKey(meta, "h1", 5)
150
+ self.assertIn("@pcp:0@dcp:0", s)
151
self.assertIn("@layer_id:5", s)
152
self.assertIn("model", s)
153
self.assertTrue(s.endswith("@h1"))
vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/metadata.py
@@ -115,7 +115,7 @@ def __hash__(self):
def to_string(self):
return (
f"{self.key_metadata.model_name}"
- f"@pcp{self.key_metadata.pcp_rank}@dcp{self.key_metadata.dcp_rank}"
+ f"@pcp:{self.key_metadata.pcp_rank}@dcp:{self.key_metadata.dcp_rank}"
f"@head_or_tp_rank:{self.key_metadata.head_or_tp_rank}"
f"@pp_rank:{self.key_metadata.pp_rank}"
f"@group:{self.key_metadata.kv_cache_group_id}"
@@ -162,7 +162,7 @@ def __hash__(self):
162
163
164
165
166
167
168
f"@cache_role:{self.key_metadata.cache_role}"
@@ -327,7 +327,7 @@ def _get_key_prefix(
327
group_metadata = self.metadata[kv_cache_group_id]
328
prefix = (
329
f"{group_metadata.model_name}"
330
- f"@pcp{group_metadata.pcp_rank}@dcp{group_metadata.dcp_rank}"
+ f"@pcp:{group_metadata.pcp_rank}@dcp:{group_metadata.dcp_rank}"
331
f"@head_or_tp_rank:{group_metadata.head_or_tp_rank}"
332
f"@pp_rank:{group_metadata.pp_rank}"
333
f"@group:{kv_cache_group_id}"
vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/pool_worker.py
@@ -2211,12 +2211,18 @@ def _replace_key_field(key: str, field: str, value: int) -> str:
2211
return f"{key[:value_start]}{value}{key[value_end:]}"
2212
2213
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.
2215
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.
2219
for pp_rank in range(self.pp_size):
- for tp_rank in range(self.get_group_tp_size(group_id)):
- for key in keys:
- tp_key = self._replace_key_field(key, "head_or_tp_rank", tp_rank)
- 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))
2226
return expanded
2227
2228
def _expand_lookup_key_variants(self, key: str, group_id: int, include_all_ranks: bool) -> list[str]:
0 commit comments