Repository navigation
[v0.27.1rc][Performance][KDA] Compose and overlap gate projections - #15978
linfeng-yuan merged 6 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) (cherry picked from commit 858b619) Signed-off-by: Dawn952 <zhaojunbo13@huawei.com>
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) (cherry picked from commit ca19043) Signed-off-by: Dawn952 <zhaojunbo13@huawei.com>
Signed-off-by: Dawn952 <zhaojunbo13@huawei.com> (cherry picked from commit 7513b78) Signed-off-by: Dawn952 <zhaojunbo13@huawei.com>
Signed-off-by: Dawn952 <zhaojunbo13@huawei.com> (cherry picked from commit 0059617) Signed-off-by: Dawn952 <zhaojunbo13@huawei.com>
Signed-off-by: linfeng-yuan <1102311262@qq.com> (cherry picked from commit da48460) Signed-off-by: Dawn952 <zhaojunbo13@huawei.com>
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 performance of the Kimi Delta Attention (KDA) mechanism by fusing gate projections and enabling concurrent execution. By composing the B, F, and G projections into a single linear layer and offloading projection computations to an auxiliary NPU stream, the implementation effectively overlaps critical path operations, resulting in improved throughput for the model. 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. Ignored Files
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:
[Ops][Feature] Fuse and overlap BFG projections in Kimi K3 Delta AttentionSuggested PR Summary:
### What this PR does / why we need it?
This PR optimizes the Kimi K3 Delta Attention (KDA) on Ascend NPUs by fusing the beta, F, and output gate (BFG) projections into a single `_KDAFusedBFGLinear` operator. It also introduces multi-stream execution (`_run_overlapped_qkv_bfg`) to overlap the BFG projection with the QKV projection's dynamic quantization and matmul, reducing execution serialization.
Feedback identifies a critical bug where weight loading fails when the model is initialized on a `meta` device. Since `self.f_a_weight` and `self.f_b_weight` are initially meta tensors, copying weights or performing premature fusion will crash. The feedback suggests materializing these parameters on the correct device during loading and deferring fusion until the main weight is loaded.
### Does this PR introduce _any_ user-facing change?
No user-facing API changes are introduced, but it improves performance for Kimi K3 models on Ascend NPUs.
### How was this patch tested?
Unit tests were added in `tests/ut/ops/test_kimi_kda.py` and updated in `tests/ut/models/test_kimi_k3_adapter.py` to verify the fused BFG linear loading, projection, overlapping stream execution, and quantization splitting.| class _KDAFusedBFGLinear(MergedColumnParallelLinear): | ||
| """Pack beta, an offline-composed F projection, and the output gate.""" | ||
|
|
||
| def __init__( | ||
| self, | ||
| hidden_size: int, | ||
| num_heads: int, | ||
| head_dim: int, | ||
| tp_size: int, | ||
| quant_config, | ||
| prefix: str, | ||
| ) -> None: | ||
| projection_size = num_heads * head_dim | ||
| super().__init__( | ||
| input_size=hidden_size, | ||
| output_sizes=[num_heads, projection_size, projection_size], | ||
| bias=False, | ||
| quant_config=quant_config, | ||
| prefix=prefix, | ||
| ) | ||
| if self.tp_size != tp_size: | ||
| raise ValueError(f"KDA fused BFG TP mismatch: layer={self.tp_size}, attention={tp_size}") | ||
| local_projection_size = projection_size // tp_size | ||
| self.f_a_weight = nn.Parameter( | ||
| self.weight.new_empty((head_dim, hidden_size)), | ||
| requires_grad=False, | ||
| ) | ||
| self.f_b_weight = nn.Parameter( | ||
| self.weight.new_empty((local_projection_size, head_dim)), | ||
| requires_grad=False, | ||
| ) | ||
| 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 | ||
|
|
||
| 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: | ||
| raise ValueError( | ||
| "KDA f_a_proj checkpoint shape mismatch: " | ||
| f"expected {tuple(param.shape)}, got {tuple(loaded_weight.shape)}" | ||
| ) | ||
| param.data.copy_(loaded_weight) | ||
| self._f_a_loaded = True | ||
| self._maybe_fuse_f_proj() | ||
|
|
||
| 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: | ||
| local_weight = loaded_weight | ||
| else: | ||
| expected_shape = (param.shape[0] * self.tp_size, param.shape[1]) | ||
| if loaded_weight.shape != expected_shape: | ||
| raise ValueError( | ||
| "KDA f_b_proj checkpoint shape mismatch: " | ||
| f"expected {expected_shape} or {tuple(param.shape)}, " | ||
| f"got {tuple(loaded_weight.shape)}" | ||
| ) | ||
| local_weight = loaded_weight.narrow( | ||
| 0, | ||
| self.tp_rank * param.shape[0], | ||
| param.shape[0], | ||
| ) | ||
| param.data.copy_(local_weight) | ||
| self._f_b_loaded = True | ||
| self._maybe_fuse_f_proj() | ||
|
|
||
| @torch.no_grad() | ||
| def _maybe_fuse_f_proj(self) -> None: | ||
| if not self._f_a_loaded or not self._f_b_loaded: | ||
| return | ||
| output_dim = getattr(self.weight, "output_dim", None) | ||
| if output_dim is None: | ||
| raise ValueError("KDA fused f_proj requires an output-sharded parameter") | ||
| shard_offset = sum(self.output_sizes[:_F_PROJ_SHARD_ID]) // self.tp_size | ||
| shard_size = self.output_sizes[_F_PROJ_SHARD_ID] // self.tp_size | ||
| param_shard = self.weight.narrow(output_dim, shard_offset, shard_size) | ||
| fused_weight = torch.matmul( | ||
| self.f_b_weight.float(), | ||
| self.f_a_weight.float(), | ||
| ).to(dtype=param_shard.dtype) | ||
| if fused_weight.shape != param_shard.shape: | ||
| raise ValueError( | ||
| "KDA composed f_proj shape mismatch: " | ||
| f"expected {tuple(param_shard.shape)}, got {tuple(fused_weight.shape)}" | ||
| ) | ||
| param_shard.copy_(fused_weight) |
There was a problem hiding this comment.
Critical Bug: Weight Loading Fails on Meta Device Initialization
When the model is initialized on a meta device (which is the default in vLLM to optimize memory usage and loading speed), self.weight is initially a meta tensor. Consequently, self.f_a_weight and self.f_b_weight are also created as meta tensors.
This causes several critical issues during weight loading:
- Meta Tensor Copy Crash: Calling
param.data.copy_(loaded_weight)onself.f_a_weightorself.f_b_weightwill raise aRuntimeErrorbecause PyTorch does not support copying to meta tensors. - Premature Fusion Crash / Silent Failure: If
f_a_projandf_b_projare loaded beforeb_projandg_proj,_maybe_fuse_f_projis called whileself.weightis still on themetadevice. This causes a crash duringself.weight.narroworparam_shard.copy_. - Missing Fusion: If we simply skip fusion when
self.weightis onmetadevice, the fused weight will never be computed or loaded becauseself.weight's default weight loader does not trigger_maybe_fuse_f_projwhen it is finally materialized.
Solution
- Materialize Meta Parameters: Check if
param.is_metaisTruein_load_f_a_weightand_load_f_b_weight, and materialize them on the correct device before copying. - Hook Weight Loader: Override
self.weight.weight_loaderto hook into the materialization and loading ofself.weight. - Defer Fusion: Skip
_maybe_fuse_f_projifself.weightis still on themetadevice, and let the overriddenself.weight.weight_loadertrigger the fusion onceself.weightis materialized and loaded.
class _KDAFusedBFGLinear(MergedColumnParallelLinear):
"""Pack beta, an offline-composed F projection, and the output gate."""
def __init__(
self,
hidden_size: int,
num_heads: int,
head_dim: int,
tp_size: int,
quant_config,
prefix: str,
) -> None:
projection_size = num_heads * head_dim
super().__init__(
input_size=hidden_size,
output_sizes=[num_heads, projection_size, projection_size],
bias=False,
quant_config=quant_config,
prefix=prefix,
)
if self.tp_size != tp_size:
raise ValueError(f"KDA fused BFG TP mismatch: layer={self.tp_size}, attention={tp_size}")
local_projection_size = projection_size // tp_size
self.f_a_weight = nn.Parameter(
self.weight.new_empty((head_dim, hidden_size)),
requires_grad=False,
)
self.f_b_weight = nn.Parameter(
self.weight.new_empty((local_projection_size, head_dim)),
requires_grad=False,
)
self.f_a_weight.weight_loader = self._load_f_a_weight
self.f_b_weight.weight_loader = self._load_f_b_weight
self._original_weight_loader = self.weight.weight_loader
self.weight.weight_loader = self._load_fused_bfg_weight
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.is_meta:
param.data = torch.empty_like(param.data, device=loaded_weight.device)
if param.shape != loaded_weight.shape:
raise ValueError(
"KDA f_a_proj checkpoint shape mismatch: "
f"expected {tuple(param.shape)}, got {tuple(loaded_weight.shape)}"
)
param.data.copy_(loaded_weight)
self._f_a_loaded = True
self._maybe_fuse_f_proj()
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.is_meta:
param.data = torch.empty_like(param.data, device=loaded_weight.device)
if loaded_weight.shape == param.shape:
local_weight = loaded_weight
else:
expected_shape = (param.shape[0] * self.tp_size, param.shape[1])
if loaded_weight.shape != expected_shape:
raise ValueError(
"KDA f_b_proj checkpoint shape mismatch: "
f"expected {expected_shape} or {tuple(param.shape)}, "
f"got {tuple(loaded_weight.shape)}"
)
local_weight = loaded_weight.narrow(
0,
self.tp_rank * param.shape[0],
param.shape[0],
)
param.data.copy_(local_weight)
self._f_b_loaded = True
self._maybe_fuse_f_proj()
def _load_fused_bfg_weight(
self,
param: nn.Parameter,
loaded_weight: torch.Tensor,
loaded_shard_id: tuple[int, ...] | int | None = None,
) -> None:
self._original_weight_loader(param, loaded_weight, loaded_shard_id)
self._maybe_fuse_f_proj()
@torch.no_grad()
def _maybe_fuse_f_proj(self) -> None:
if not self._f_a_loaded or not self._f_b_loaded or self.weight.is_meta:
return
output_dim = getattr(self.weight, "output_dim", None)
if output_dim is None:
raise ValueError("KDA fused f_proj requires an output-sharded parameter")
shard_offset = sum(self.output_sizes[:_F_PROJ_SHARD_ID]) // self.tp_size
shard_size = self.output_sizes[_F_PROJ_SHARD_ID] // self.tp_size
param_shard = self.weight.narrow(output_dim, shard_offset, shard_size)
fused_weight = torch.matmul(
self.f_b_weight.float(),
self.f_a_weight.float(),
).to(dtype=param_shard.dtype)
if fused_weight.shape != param_shard.shape:
raise ValueError(
"KDA composed f_proj shape mismatch: "
f"expected {tuple(param_shard.shape)}, got {tuple(fused_weight.shape)}"
)
param_shard.copy_(fused_weight)Signed-off-by: Dawn952 <zhaojunbo13@huawei.com>
cc3537b
into
vllm-project:releases/v0.27.1rc
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.1AttentionSelectorConfigconstructor.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 source-weight reloads;scale_algwhen DynamicQuant is split from the linear method;in_proj_qkvgfabpath unchanged.ACL graph and stream ordering
The auxiliary stream records
bfg_readyonly 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), thenbfg_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, and0059617c5d9e85c4ee2185110ee2cec950698533, plus the v0.27.1 CI-selection commitda4846002da04346772c12774418bc2cc8d40ed4and the release-onlyAttentionSelectorConfigcompatibility 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.1rccompleted successfully (26 checks passed; 3 intentionally skipped).vLLM main: vllm-project/vllm@ba07e4a