Repository navigation
[Performance][KDA] Compose and overlap gate projections on main - #15416
linfeng-yuan merged 5 commits into
Conversation
Port vllm-project#14497 onto the Kimi K3 implementation now available on main. Compose the full-rank F projection at load time, overlap the float BFG path with MXFP QKV preprocessing, and keep the auxiliary stream fully joined for ACL graph capture. Preserve the configured W8A8 MXFP8 dynamic quantization scale algorithm. Signed-off-by: Dawn952 <zhaojunbo13@huawei.com> (cherry picked from commit 37df0a2)
Move fused BFG splitting and beta FP32 sigmoid onto the auxiliary stream after QKV matmul is enqueued. Mark mixed-path beta as preprocessed so the v0.27 dispatch path only slices it, while ordinary raw beta retains its existing sigmoid contract. Keep the graph-safe auxiliary tail join and make the stream-order test mypy-safe. Signed-off-by: Dawn952 <zhaojunbo13@huawei.com> (cherry picked from commit fd8125e)
Summary of ChangesHello, I'm Gemini Code Assist1! I'm currently reviewing this pull request and will post my feedback shortly. In the meantime, here's a summary to help you and other reviewers quickly get up to speed! This pull request optimizes the Kimi K3 Delta Attention (KDA) performance by introducing a fused BFG projection and a two-stage auxiliary stream schedule. By packing beta, F projection, and output gate operations and overlapping them with main-stream computations like MXFP DynamicQuant and QKV GEMM, the implementation reduces latency and improves hardware utilization on NPU devices. Highlights
New Features🧠 You can now enable Memory (public preview) to help Gemini Code Assist learn from your team's feedback. This makes future code reviews more consistent and personalized to your project's style. Click here to enable Memory in your admin console. Using Gemini Code AssistThe full guide for Gemini Code Assist can be found on our documentation page, here are some quick tips. Invoking Gemini You can request assistance from Gemini at any point by creating a comment using either
Customization To customize the Gemini Code Assist for GitHub experience, repository maintainers can create a configuration file and/or provide a custom code review style guide (such as PEP-8 for Python) by creating and adding files to a Limitations & Feedback Gemini Code Assist may make mistakes. Please leave feedback on any instances where its feedback is incorrect or counterproductive. You can react with 👍 and 👎 on @gemini-code-assist comments. If you're interested in giving your feedback about your experience with Gemini Code Assist for GitHub and other Google products, sign up here. Footnotes
|
|
👋 Hi! Thank you for contributing to the vLLM Ascend project. The following points will speed up your PR merge:
If CI fails, you can run linting and testing checks locally according Contributing and Testing. |
There was a problem hiding this comment.
Code Review
Suggested PR Title:\n\nmarkdown\n[Ops][Feature] Optimize Kimi K3 Delta Attention with Fused BFG Projection and Multi-Stream Overlap\n\n\nSuggested PR Summary:\n\nmarkdown\n### What this PR does / why we need it?\nThis pull request optimizes the Kimi K3 Delta Attention (KDA) implementation on Ascend NPU. It splits the precision groups for QKV and BFG projections, introducing a fused BFG linear layer (`_KDAFusedBFGLinear`) that offline-composes the F projection and packs it with beta and the output gate. It also implements multi-stream execution (`_run_overlapped_qkv_bfg`) to overlap the dynamic quantization of QKV with the BFG projection, improving hardware utilization.\n\nFeedback on the implementation highlights critical issues when initializing models on the `meta` device. Specifically, loading weights when parameters are on the `meta` device will raise `RuntimeError` during copy operations, and the fusion logic (`_maybe_fuse_f_proj`) can be bypassed or fail if weights are loaded out of order.\n\n### Does this PR introduce _any_ user-facing change?\nNo, this is an internal performance optimization and refactoring of the Kimi K3 attention layer.\n\n### How was this patch tested?\nThe changes are covered by new unit tests in `tests/ut/models/test_kimi_k3_adapter.py` and `tests/ut/ops/test_kimi_kda.py` verifying weight loading, fused projection correctness, stream synchronization, and quantization integration.\n
| self.f_a_weight.weight_loader = self._load_f_a_weight | ||
| self.f_b_weight.weight_loader = self._load_f_b_weight | ||
| self._f_a_loaded = False | ||
| self._f_b_loaded = False |
There was a problem hiding this comment.
If the checkpoint weights are loaded in an order where f_a_proj and f_b_proj are loaded before self.weight, _maybe_fuse_f_proj will be called while self.weight is still on the meta device, causing a copy failure or skipping fusion without a way to re-trigger it.
To prevent this, wrap self.weight.weight_loader to automatically trigger _maybe_fuse_f_proj once self.weight is loaded and moved off the meta device.
self.f_a_weight.weight_loader = self._load_f_a_weight
self.f_b_weight.weight_loader = self._load_f_b_weight
original_weight_loader = self.weight.weight_loader
def wrapped_weight_loader(*args, **kwargs):
original_weight_loader(*args, **kwargs)
self._maybe_fuse_f_proj()
self.weight.weight_loader = wrapped_weight_loader
self._f_a_loaded = False
self._f_b_loaded = False| def _load_f_a_weight( | ||
| self, | ||
| param: nn.Parameter, | ||
| loaded_weight: torch.Tensor, | ||
| loaded_shard_id: tuple[int, ...] | int | None = None, | ||
| ) -> None: | ||
| del loaded_shard_id | ||
| if param.shape != loaded_weight.shape: |
There was a problem hiding this comment.
When the model is initialized on the meta device, self.f_a_weight is also on the meta device. Calling param.data.copy_(loaded_weight) directly on a meta tensor will raise a RuntimeError.
We should check if the parameter is on the meta device and allocate it on the correct device (either self.weight.device if it is already loaded, or CPU/NPU) before copying.
def _load_f_a_weight(
self,
param: nn.Parameter,
loaded_weight: torch.Tensor,
loaded_shard_id: tuple[int, ...] | int | None = None,
) -> None:
del loaded_shard_id
if param.device == torch.device("meta"):
device = self.weight.device if self.weight.device != torch.device("meta") else "cpu"
param.data = torch.empty(param.shape, device=device, dtype=param.dtype)
if param.shape != loaded_weight.shape:| def _load_f_b_weight( | ||
| self, | ||
| param: nn.Parameter, | ||
| loaded_weight: torch.Tensor, | ||
| loaded_shard_id: tuple[int, ...] | int | None = None, | ||
| ) -> None: | ||
| del loaded_shard_id | ||
| if loaded_weight.shape == param.shape: |
There was a problem hiding this comment.
When the model is initialized on the meta device, self.f_b_weight is also on the meta device. Calling param.data.copy_(local_weight) directly on a meta tensor will raise a RuntimeError.
We should check if the parameter is on the meta device and allocate it on the correct device before copying.
| def _load_f_b_weight( | |
| self, | |
| param: nn.Parameter, | |
| loaded_weight: torch.Tensor, | |
| loaded_shard_id: tuple[int, ...] | int | None = None, | |
| ) -> None: | |
| del loaded_shard_id | |
| if loaded_weight.shape == param.shape: | |
| def _load_f_b_weight( | |
| self, | |
| param: nn.Parameter, | |
| loaded_weight: torch.Tensor, | |
| loaded_shard_id: tuple[int, ...] | int | None = None, | |
| ) -> None: | |
| del loaded_shard_id | |
| if param.device == torch.device("meta"): | |
| device = self.weight.device if self.weight.device != torch.device("meta") else "cpu" | |
| param.data = torch.empty(param.shape, device=device, dtype=param.dtype) | |
| if loaded_weight.shape == param.shape: |
| @torch.no_grad() | ||
| def _maybe_fuse_f_proj(self) -> None: | ||
| if not self._f_a_loaded or not self._f_b_loaded: | ||
| return |
There was a problem hiding this comment.
If self.weight is still on the meta device when _maybe_fuse_f_proj is called, we should skip the fusion and wait for self.weight to be loaded (which will trigger the fusion via the wrapped weight loader).
| @torch.no_grad() | |
| def _maybe_fuse_f_proj(self) -> None: | |
| if not self._f_a_loaded or not self._f_b_loaded: | |
| return | |
| @torch.no_grad() | |
| def _maybe_fuse_f_proj(self) -> None: | |
| if not self._f_a_loaded or not self._f_b_loaded: | |
| return | |
| if self.weight.device == torch.device("meta"): | |
| return |
Signed-off-by: Dawn952 <zhaojunbo13@huawei.com>
Signed-off-by: Dawn952 <zhaojunbo13@huawei.com>
…-project#15416) ### What this PR does / why we need it? This recreates vllm-project#15168 without the unrelated `mla_v1.py` head-padding change and its corresponding test changes. For the existing mixed-precision Kimi K3 KDA layout, this change: - composes `f_proj.weight = f_b_proj.weight @ f_a_proj.weight` after checkpoint loading and later source-weight reloads; - packs beta, the composed F projection, and the output gate into one floating-point BFG projection; - runs a two-stage auxiliary-stream schedule: 1. main-stream MXFP DynamicQuant overlaps the auxiliary BFG GEMM; 2. main-stream QKV GEMM overlaps auxiliary B/F/G split, beta FP32 sigmoid, and gate reshaping; - marks mixed-path beta as preprocessed so the v0.27 dispatch path only slices it, while ordinary upstream raw beta retains its existing FP32 sigmoid path; - preserves the configured W8A8 MXFP8 `scale_alg` when DynamicQuant is split from the linear method; - routes global or local F shards through the v0.27 packed loader; - leaves the ordinary upstream `in_proj_qkvgfab` path unchanged. Compared with vllm-project#15168, this PR does not modify: - `vllm_ascend/attention/mla_v1.py`; - `tests/ut/attention/a2/test_mla_v1.py`. ### ACL graph and stream-order assessment The auxiliary stream records `bfg_ready` as its final node, and the main stream waits for that exact tail event. No auxiliary operation is queued after the join point. The four-event schedule is intentional: 1. `hidden_states_ready` forks BFG from main; 2. `bfg_projection_ready` serializes the BFG and QKV Cube matmuls; 3. `quant_ready` starts auxiliary beta vector work only after main DynamicQuant; 4. `bfg_ready` joins the complete auxiliary tail before KDA consumes its outputs. The live-token slices for `mixed_qkv`, `g1`, and `g2` remain at the eager dispatch boundary. They are host-side basic views: after QKV GEMM and the auxiliary tail wait are enqueued, Python creates those views while the asynchronous device work is still running. Moving `num_actual_tokens` into the projection path would not add an NPU kernel overlap and would weaken the runtime metadata boundary. Static assessment: the stream is fully joined for ACLGraph capture and mixed beta is transformed exactly once. A real NPU ACLGraph capture/replay has not been run for this main port. ### How was this patch tested? No additional local or NPU validation was run for this replacement PR. The retained KDA implementation and focused tests are unchanged from vllm-project#15168; this replacement only excludes the MLA implementation and test diffs. ### Does this PR introduce any user-facing change? No API or configuration change. Runtime behavior changes only for Kimi K3 KDA layers using the existing mixed-precision full-rank gate layout. - vLLM main: vllm-project/vllm@ba07e4a --------- Signed-off-by: Dawn952 <zhaojunbo13@huawei.com>
|
/cherry-pick releases/v0.27.1rc |
…15978) ### What this PR does / why we need it? Backport #15416 to `releases/v0.27.1rc`, retaining its Kimi K3 KDA mixed-precision projection schedule and v0.27-compatible dispatch. The backport also includes the release branch CI-selection commit and the compatibility fix for the v0.27.1 `AttentionSelectorConfig` constructor. For the existing mixed-precision Kimi K3 KDA layout, this change: - composes `f_proj.weight = f_b_proj.weight @ f_a_proj.weight` after checkpoint loading and source-weight reloads; - packs beta, the composed F projection, and the output gate into a floating-point BFG projection; - overlaps main-stream MXFP DynamicQuant with the auxiliary BFG GEMM, then overlaps main-stream QKV GEMM with auxiliary B/F/G split, beta FP32 sigmoid, and gate reshape; - marks mixed-path beta as preprocessed so the v0.27 dispatch path slices it, while ordinary raw beta retains its existing FP32 sigmoid path; - preserves the configured W8A8 MXFP8 `scale_alg` when DynamicQuant is split from the linear method; - supports both global and local F shards through the v0.27 packed loader; - leaves the ordinary upstream `in_proj_qkvgfab` path unchanged. ### ACL graph and stream ordering The auxiliary stream records `bfg_ready` only after its complete BFG tail. The main stream waits for that event before KDA consumes the outputs, so the auxiliary stream is joined for ACLGraph capture. The intended order is `hidden_states_ready` (fork BFG), `bfg_projection_ready` (serialize BFG/QKV Cube matmuls), `quant_ready` (start beta vector work after DynamicQuant), then `bfg_ready` (join the auxiliary tail). Live-token slicing remains at the eager dispatch boundary and does not add an NPU kernel dependency. ### Release-branch differences This is a manual backport of #15416 commits `858b61937e81de059ac6ed593fb7aec2941f6f47`, `ca19043f2ef7ce798a18982f29970a2aa6e18d1e`, `7513b78278e4aa1f43da6a9082f18f50884235da`, and `0059617c5d9e85c4ee2185110ee2cec950698533`, plus the v0.27.1 CI-selection commit `da4846002da04346772c12774418bc2cc8d40ed4` and the release-only `AttentionSelectorConfig` compatibility fix. ### Does this PR introduce any user-facing change? No API or configuration change. Runtime behavior changes only for Kimi K3 KDA layers using the existing mixed-precision full-rank gate layout. ### How was this patch tested? - Focused Kimi K3 adapter and KDA unit coverage is included in the backport. - PR CI on `releases/v0.27.1rc` completed successfully (26 checks passed; 3 intentionally skipped). - vLLM main: vllm-project/vllm@ba07e4a --------- Signed-off-by: Dawn952 <zhaojunbo13@huawei.com> Signed-off-by: linfeng-yuan <1102311262@qq.com> Co-authored-by: linfeng-yuan <1102311262@qq.com>
…-project#15416) ### What this PR does / why we need it? This recreates vllm-project#15168 without the unrelated `mla_v1.py` head-padding change and its corresponding test changes. For the existing mixed-precision Kimi K3 KDA layout, this change: - composes `f_proj.weight = f_b_proj.weight @ f_a_proj.weight` after checkpoint loading and later source-weight reloads; - packs beta, the composed F projection, and the output gate into one floating-point BFG projection; - runs a two-stage auxiliary-stream schedule: 1. main-stream MXFP DynamicQuant overlaps the auxiliary BFG GEMM; 2. main-stream QKV GEMM overlaps auxiliary B/F/G split, beta FP32 sigmoid, and gate reshaping; - marks mixed-path beta as preprocessed so the v0.27 dispatch path only slices it, while ordinary upstream raw beta retains its existing FP32 sigmoid path; - preserves the configured W8A8 MXFP8 `scale_alg` when DynamicQuant is split from the linear method; - routes global or local F shards through the v0.27 packed loader; - leaves the ordinary upstream `in_proj_qkvgfab` path unchanged. Compared with vllm-project#15168, this PR does not modify: - `vllm_ascend/attention/mla_v1.py`; - `tests/ut/attention/a2/test_mla_v1.py`. ### ACL graph and stream-order assessment The auxiliary stream records `bfg_ready` as its final node, and the main stream waits for that exact tail event. No auxiliary operation is queued after the join point. The four-event schedule is intentional: 1. `hidden_states_ready` forks BFG from main; 2. `bfg_projection_ready` serializes the BFG and QKV Cube matmuls; 3. `quant_ready` starts auxiliary beta vector work only after main DynamicQuant; 4. `bfg_ready` joins the complete auxiliary tail before KDA consumes its outputs. The live-token slices for `mixed_qkv`, `g1`, and `g2` remain at the eager dispatch boundary. They are host-side basic views: after QKV GEMM and the auxiliary tail wait are enqueued, Python creates those views while the asynchronous device work is still running. Moving `num_actual_tokens` into the projection path would not add an NPU kernel overlap and would weaken the runtime metadata boundary. Static assessment: the stream is fully joined for ACLGraph capture and mixed beta is transformed exactly once. A real NPU ACLGraph capture/replay has not been run for this main port. ### How was this patch tested? No additional local or NPU validation was run for this replacement PR. The retained KDA implementation and focused tests are unchanged from vllm-project#15168; this replacement only excludes the MLA implementation and test diffs. ### Does this PR introduce any user-facing change? No API or configuration change. Runtime behavior changes only for Kimi K3 KDA layers using the existing mixed-precision full-rank gate layout. - vLLM main: vllm-project/vllm@ba07e4a --------- Signed-off-by: Dawn952 <zhaojunbo13@huawei.com>
What this PR does / why we need it?
This recreates #15168 without the unrelated
mla_v1.pyhead-padding change and its corresponding test changes.For the existing mixed-precision Kimi K3 KDA layout, this change:
f_proj.weight = f_b_proj.weight @ f_a_proj.weightafter checkpoint loading and later source-weight reloads;scale_algwhen DynamicQuant is split from the linear method;in_proj_qkvgfabpath unchanged.Compared with #15168, this PR does not modify:
vllm_ascend/attention/mla_v1.py;tests/ut/attention/a2/test_mla_v1.py.ACL graph and stream-order assessment
The auxiliary stream records
bfg_readyas its final node, and the main stream waits for that exact tail event. No auxiliary operation is queued after the join point.The four-event schedule is intentional:
hidden_states_readyforks BFG from main;bfg_projection_readyserializes the BFG and QKV Cube matmuls;quant_readystarts auxiliary beta vector work only after main DynamicQuant;bfg_readyjoins the complete auxiliary tail before KDA consumes its outputs.The live-token slices for
mixed_qkv,g1, andg2remain at the eager dispatch boundary. They are host-side basic views: after QKV GEMM and the auxiliary tail wait are enqueued, Python creates those views while the asynchronous device work is still running. Movingnum_actual_tokensinto the projection path would not add an NPU kernel overlap and would weaken the runtime metadata boundary.Static assessment: the stream is fully joined for ACLGraph capture and mixed beta is transformed exactly once. A real NPU ACLGraph capture/replay has not been run for this main port.
How was this patch tested?
No additional local or NPU validation was run for this replacement PR. The retained KDA implementation and focused tests are unchanged from #15168; this replacement only excludes the MLA implementation and test diffs.
Does this PR introduce any user-facing change?
No API or configuration change. Runtime behavior changes only for Kimi K3 KDA layers using the existing mixed-precision full-rank gate layout.