-
Notifications
You must be signed in to change notification settings - Fork 2.4k
[BugFix][v0.23.0][KV Pool] Include MTP KV in layerwise AscendStore transfer #13454
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
yiz-liu
merged 1 commit into
vllm-project:releases/v0.23.0
from
tyy0829:rebase-pr-13384-v0.23.0
Aug 4, 2026
Merged
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
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -275,38 +275,10 @@ def _init_state_vars(self) -> None: | |
| self._allocated_gvas: dict[str, int] = {} | ||
|
|
||
| def _init_layerwise_config(self) -> None: | ||
| # Build mapping: physical_layer -> [(group_id, layer_idx_in_group), ...] | ||
| # layer_idx_in_group is the index of the physical layer within the | ||
| # group (not the index in layer_names). Multiple layer_names at the | ||
| # same physical layer (e.g. indexer.k_cache + attn) are treated as | ||
| # multiple cache tensors of ONE layer (caches_per_layer > 1). | ||
| self.physical_layer_to_group_layers: dict[int, list[tuple[int, int]]] = {} | ||
|
|
||
| if self.kv_cache_config is not None and self.num_kv_cache_groups > 1: | ||
| for group_id, group_spec in enumerate(self.kv_cache_config.kv_cache_groups): | ||
| # Map each unique physical layer to a sequential layer_idx_in_group | ||
| phys_to_layer_idx: dict[int, int] = {} | ||
| for layer_name in group_spec.layer_names: | ||
| physical_layer = self._extract_physical_layer_index(layer_name) | ||
| if physical_layer >= getattr(self.hf_config, "num_hidden_layers", self.num_layers): | ||
| continue | ||
| if physical_layer not in phys_to_layer_idx: | ||
| phys_to_layer_idx[physical_layer] = len(phys_to_layer_idx) | ||
|
|
||
| # Add one entry per unique physical layer (no duplicates) | ||
| for physical_layer, layer_idx_in_group in phys_to_layer_idx.items(): | ||
| existing = self.physical_layer_to_group_layers.setdefault(physical_layer, []) | ||
| entry = (group_id, layer_idx_in_group) | ||
| if entry not in existing: | ||
| existing.append(entry) | ||
|
|
||
| logger.info( | ||
| "layerwise group %d: %d layer_names, %d unique physical layers, caches_per_layer=%d", | ||
| group_id, | ||
| len(group_spec.layer_names), | ||
| len(phys_to_layer_idx), | ||
| len(group_spec.layer_names) // max(1, len(phys_to_layer_idx)), | ||
| ) | ||
| # Registration provides the authoritative layer layout. In | ||
| # particular, model_config.get_num_layers() only reports the local PP | ||
| # target layers and does not include registered MTP draft layers. | ||
| self.local_layer_to_group_layers: dict[int, list[tuple[int, int]]] = {} | ||
|
|
||
| self.layer_load_tasks: list[list[LayerTransferTask]] = [[] for _ in range(self.num_layers)] | ||
| self.layer_save_tasks: list[list[LayerTransferTask]] = [[] for _ in range(self.num_layers)] | ||
|
|
@@ -318,10 +290,9 @@ def _init_layerwise_config(self) -> None: | |
| self.sync_save_events: list[torch.npu.Event] | None = None | ||
|
|
||
| logger.info( | ||
| "layerwise config: num_layers=%d num_groups=%d physical_layer_to_group_layers_sample=%s", | ||
| "layerwise config: num_layers=%d num_groups=%d", | ||
| self.num_layers, | ||
| self.num_kv_cache_groups, | ||
| {k: v for k, v in list(self.physical_layer_to_group_layers.items())[:3]}, | ||
| ) | ||
|
|
||
| def _build_group_layer_builders(self) -> list[LayerBatchBuilder]: | ||
|
|
@@ -613,8 +584,6 @@ def _infer_cache_group_metadata(self, group_id: int, layer_names: list[str]): | |
| physical_layers = set() | ||
| for layer_name in layer_names: | ||
| phys = self._extract_physical_layer_index(layer_name) | ||
| if phys >= getattr(self.hf_config, "num_hidden_layers", self.num_layers): | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I think the draft layers will not be discarded now. |
||
| continue | ||
| physical_layers.add(phys) | ||
| cache_or_caches = self.kv_caches[layer_name] | ||
| for cache in self._as_cache_tuple(cache_or_caches): | ||
|
|
@@ -628,6 +597,38 @@ def _infer_cache_group_metadata(self, group_id: int, layer_names: list[str]): | |
| self.group_block_stride[group_id] = group_block_strides | ||
| self.group_num_layers[group_id] = len(physical_layers) | ||
|
|
||
| def _configure_registered_layerwise_layers(self, cache_groups: list[tuple[int, list[str]]]) -> None: | ||
| """Map registered physical layers to dense local and group indices.""" | ||
| groups_by_physical: dict[int, list[tuple[int, int]]] = {} | ||
|
|
||
| for group_id, layer_names in cache_groups: | ||
| group_physical_layers: list[int] = [] | ||
| seen_in_group: set[int] = set() | ||
| for layer_name in layer_names: | ||
| physical_layer = self._extract_physical_layer_index(layer_name) | ||
| if physical_layer not in seen_in_group: | ||
| seen_in_group.add(physical_layer) | ||
| group_physical_layers.append(physical_layer) | ||
|
|
||
| for layer_idx_in_group, physical_layer in enumerate(group_physical_layers): | ||
| groups_by_physical.setdefault(physical_layer, []).append((group_id, layer_idx_in_group)) | ||
|
|
||
| physical_layers = sorted(groups_by_physical) | ||
| self.local_layer_to_group_layers = { | ||
| local_layer: groups_by_physical[physical_layer] | ||
| for local_layer, physical_layer in enumerate(physical_layers) | ||
| } | ||
|
|
||
| original_num_layers = self.num_layers | ||
| self.num_layers = len(physical_layers) | ||
| self.layer_load_tasks = [[] for _ in range(self.num_layers)] | ||
| self.layer_save_tasks = [[] for _ in range(self.num_layers)] | ||
| logger.info( | ||
| "KVPoolWorker: configured %d registered layers (was %d).", | ||
| self.num_layers, | ||
| original_num_layers, | ||
| ) | ||
|
|
||
| def _align_kv_ptrs(self, registered_regions: dict[int, tuple[int, int]]): | ||
| """ | ||
| In hybrid scenario, where a KVCacheTensor is shared by multiple layers, | ||
|
|
@@ -704,27 +705,17 @@ def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): | |
| ptrs = [start for start, _ in registered_regions.values()] | ||
| lengths = [end - start for start, end in registered_regions.values()] | ||
|
|
||
| cache_groups: list[tuple[int, list[str]]] | ||
| if self.kv_cache_config is not None and self.use_hybrid: | ||
| for group_id, group_spec in enumerate(self.kv_cache_config.kv_cache_groups): | ||
| self._infer_cache_group_metadata(group_id, group_spec.layer_names) | ||
| cache_groups = [ | ||
| (group_id, [layer_name for layer_name in group_spec.layer_names if layer_name in kv_caches]) | ||
| for group_id, group_spec in enumerate(self.kv_cache_config.kv_cache_groups) | ||
| ] | ||
| else: | ||
| self._infer_cache_group_metadata(0, list(kv_caches.keys())) | ||
|
|
||
| # group_num_layers is computed from the actual kv_caches dict which | ||
| # includes ALL attention layers (main + MTP). For single-group models, | ||
| # sum(group_num_layers.values()) equals the physical layer count | ||
| # (including MTP). For multi-group models, it counts (group, layer) | ||
| # pairs which is NOT the physical layer count — keep the original | ||
| # num_layers (physical layers) in that case. | ||
| original_num_layers = self.num_layers | ||
| new_num_layers = sum(self.group_num_layers.values()) | ||
| if self.num_kv_cache_groups == 1 and new_num_layers != original_num_layers: | ||
| self.num_layers = new_num_layers | ||
| logger.info( | ||
| "KVPoolWorker: updated num_layers %d -> %d (includes MTP/spec-decode draft layers).", | ||
| original_num_layers, | ||
| self.num_layers, | ||
| ) | ||
| cache_groups = [(0, list(kv_caches.keys()))] | ||
| for group_id, layer_names in cache_groups: | ||
| self._infer_cache_group_metadata(group_id, layer_names) | ||
| self._configure_registered_layerwise_layers(cache_groups) | ||
|
|
||
| self.page_size_bytes = sum(self.block_len) | ||
| self.token_database.set_group_buffers( | ||
|
|
@@ -1344,17 +1335,17 @@ def _build_shared_load_data(self) -> None: | |
| def process_layer_data(self, requests: list[ReqMeta]) -> None: | ||
| if not requests: | ||
| return | ||
| for physical_layer in range(self.num_layers): | ||
| group_layers = self.physical_layer_to_group_layers.get(physical_layer, [(0, physical_layer)]) | ||
| for local_layer in range(self.num_layers): | ||
| group_layers = self.local_layer_to_group_layers.get(local_layer, [(0, local_layer)]) | ||
| for group_id, layer_idx_in_group in group_layers: | ||
| self._process_save_for_layer_batch(requests, physical_layer, group_id, layer_idx_in_group) | ||
| self._process_save_for_layer_batch(requests, local_layer, group_id, layer_idx_in_group) | ||
| self._alloc_gvas_for_save(requests) | ||
| self._build_shared_save_data() | ||
| self._prepare_load_gvas(requests) | ||
| for physical_layer in range(self.num_layers): | ||
| group_layers = self.physical_layer_to_group_layers.get(physical_layer, [(0, physical_layer)]) | ||
| for local_layer in range(self.num_layers): | ||
| group_layers = self.local_layer_to_group_layers.get(local_layer, [(0, local_layer)]) | ||
| for group_id, layer_idx_in_group in group_layers: | ||
| self._process_load_for_layer_batch(requests, physical_layer, group_id, layer_idx_in_group) | ||
| self._process_load_for_layer_batch(requests, local_layer, group_id, layer_idx_in_group) | ||
| self._build_shared_load_data() | ||
|
|
||
| def _submit_ready_layer_loads(self) -> None: | ||
|
|
||
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
The Pull Request title and summary do not adhere to the repository's style guide. Please update the PR title and summary to match the following suggested formats:\n\nSuggested PR Title:\n\n
markdown\n[releases/v0.23.0][Attention][BugFix] transfer registered MTP layers\n\n\nSuggested PR Summary:\n\nmarkdown\n### What this PR does / why we need it?\n\nThis PR builds layerwise execution and group offsets from the KV caches registered on each worker. This includes MTP caches while keeping PP-local task indices and dense per-group storage indices for hybrid models.\n\n### Does this PR introduce _any_ user-facing change?\n\nNo.\n\n### How was this patch tested?\n\nTested with new unit tests in `tests/ut/distributed/ascend_store/test_pool_worker.py`:\n- `test_registered_layer_layout_includes_mtp_multi_group_and_pp`\n- `test_lookup_reuses_grouped_hashes_for_hit_resolution`\nReferences