Skip to content

[v0.27.1rc][Performance][KDA] Compose and overlap gate projections - #15978

Merged
linfeng-yuan merged 6 commits into
vllm-project:releases/v0.27.1rcfrom
Dawn952:backport-v0.27.1-kda-projection-overlap
Sep 8, 2026
Merged

linfeng-yuan merged 6 commits into
vllm-project:releases/v0.27.1rcfrom
Dawn952:backport-v0.27.1-kda-projection-overlap

Conversation

@Dawn952

@Dawn952 Dawn952 commented Sep 7, 2026 •

Copy link
Copy Markdown
Contributor

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

Dawn952 and others added 5 commits September 7, 2026 23:28
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>
@gemini-code-assist

Copy link
Copy Markdown
Contributor

Summary of Changes

Hello, 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

  • Fused BFG Linear Layer: Introduced _KDAFusedBFGLinear to handle the composition and packing of Beta, F-projection, and G-projection weights, optimizing memory and compute usage.
  • Multi-stream Overlapping: Implemented multi-stream execution using an auxiliary NPU stream to overlap QKV quantization with BFG projection computations, reducing overall latency.
  • Attention Logic Update: Refactored AscendKimiK3DeltaAttention to utilize the new fused projection and stream-switching logic during the forward pass.
  • Test Coverage: Added and updated unit tests to verify the correctness of the fused linear layer, checkpoint mapping, and the multi-stream overlapping execution flow.
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
  • Ignored by pattern: .github/workflows/** (1)
    • .github/workflows/pr_test.yaml
Using Gemini Code Assist

The 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 /gemini <command> or @gemini-code-assist <command>. Below is a summary of the supported commands on the current page.

Feature Command Description
Code Review /gemini review Performs a code review for the current pull request in its current state.
Pull Request Summary /gemini summary Provides a summary of the current pull request in its current state.
Comment @gemini-code-assist Responds in comments when explicitly tagged, both in pull request comments and review comments.
Help /gemini help Displays a list of available commands.

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 .gemini/ folder in the base of the repository. Detailed instructions can be found here.

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

  1. Review the Privacy Notices, Generative AI Prohibited Use Policy, Terms of Service, and learn how to configure Gemini Code Assist in GitHub here. Gemini can make mistakes, so double check it and use code with caution. ↩

@github-actions

github-actions Bot commented Sep 7, 2026

Copy link
Copy Markdown
Contributor

👋 Hi! Thank you for contributing to the vLLM Ascend project. The following points will speed up your PR merge:‌‌

  • A PR should do only one thing, smaller PRs enable faster reviews.
  • Every PR should include unit tests and end-to-end tests ‌to ensure it works and is not broken by other future PRs.
  • Write the commit message by fulfilling the PR description to help reviewer and future developers understand.

If CI fails, you can run linting and testing checks locally according Contributing and Testing.

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Code Review

Suggested PR Title:

[Ops][Feature] Fuse and overlap BFG projections in Kimi K3 Delta Attention

Suggested 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.

Comment on lines +53 to +150
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)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

critical

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:

  1. Meta Tensor Copy Crash: Calling param.data.copy_(loaded_weight) on self.f_a_weight or self.f_b_weight will raise a RuntimeError because PyTorch does not support copying to meta tensors.
  2. Premature Fusion Crash / Silent Failure: If f_a_proj and f_b_proj are loaded before b_proj and g_proj, _maybe_fuse_f_proj is called while self.weight is still on the meta device. This causes a crash during self.weight.narrow or param_shard.copy_.
  3. Missing Fusion: If we simply skip fusion when self.weight is on meta device, the fused weight will never be computed or loaded because self.weight's default weight loader does not trigger _maybe_fuse_f_proj when it is finally materialized.

Solution

  • Materialize Meta Parameters: Check if param.is_meta is True in _load_f_a_weight and _load_f_b_weight, and materialize them on the correct device before copying.
  • Hook Weight Loader: Override self.weight.weight_loader to hook into the materialization and loading of self.weight.
  • Defer Fusion: Skip _maybe_fuse_f_proj if self.weight is still on the meta device, and let the overridden self.weight.weight_loader trigger the fusion once self.weight is 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)

@linfeng-yuan linfeng-yuan added the ready-precise run selected e2e test for pr label Sep 7, 2026
Signed-off-by: Dawn952 <zhaojunbo13@huawei.com>
@linfeng-yuan
linfeng-yuan merged commit cc3537b into vllm-project:releases/v0.27.1rc Sep 8, 2026
29 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants