# FlashKDA Two-Kernel Fusion Strategy: Why Splitting K1 and K2 Beats a Single CUDA Kernel

> Discover how FlashKDA's two-kernel fusion strategy (K1 and K2) boosts performance by 15% at least, eliminating idle SMs and outperforming monolithic kernels for Kimi Delta Attention.

- Repository: [Moonshot AI/FlashKDA](https://github.com/MoonshotAI/FlashKDA)
- Tags: deep-dive
- Published: 2026-07-31

---

**FlashKDA's two-kernel fusion strategy decomposes Kimi Delta Attention into separate token-parallel (K1) and head-parallel (K2) CUDA kernels to eliminate idle SMs and achieve at least 15% end-to-end speedup over a monolithic kernel approach.**

The MoonshotAI/FlashKDA library implements a novel two-kernel fusion strategy that splits the Kimi Delta Attention computation into distinct preparation and recurrence stages. Unlike monolithic kernel designs that force uniform parallelism across heterogeneous workloads, this architecture assigns token-parallel work to K1 and head-parallel work to K2. By examining the dispatch logic in [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp) and the kernel implementations, we can see how this separation maximizes GPU occupancy and throughput.

## Architecture of the Two-Kernel Fusion Strategy

### Kernel 1 (K1): Token-Parallel Preparation

In `csrc/smxx/fwd_kernel1.cuh`, the K1 kernel executes with a grid configuration of `N × H × num_chunks`, enabling massive token-level parallelism across the sequence. This stage handles the **activation of the gate *g***, **L2-normalisation**, **decay calculations**, construction of the **`L`** and **`Mqk`** matrices, and **matrix inversion**. Because the grid scales with the number of chunks, K1 can saturate thousands of Streaming Multiprocessors (SMs) simultaneously without being constrained by subsequent operations.

### Kernel 2 (K2): Head-Parallel Recurrence

The K2 kernel, defined in `csrc/smxx/fwd_kernel2.cuh`, operates with a more focused grid of `N × H`, processing the **chunk-by-chunk recurrence** using the delta-rule, **output projection**, and **accumulation of the running state**. This head-parallel design aligns with the inherently sequential nature of the recurrence computation, which exhibits lower parallelism than the token-level preparation work. The shared-memory and TMA plumbing for K2 is managed in `csrc/smxx/fwd_launch.cu`, which handles pipelines and state conversion between stages.

## Why a Single Fused Kernel Creates Bottlenecks

A **single fused kernel** forces the entire pipeline to run with the lowest common denominator of parallelism. In a monolithic design, the token-parallel work of K1—which can utilize thousands of SMs—would be throttled by the **much lower parallelism** of the K2 recurrence. This mismatch leaves many SMs idle during the recurrence phase, limiting overall throughput and wasting GPU resources.

By splitting the pipeline into two distinct launches:
- **K1** exploits massive token-parallelism independently, saturating GPU compute resources.
- **K2** runs in its own launch with a grid sized exactly to its needs (`N × H`), ensuring efficient execution without starving K1 of parallelism.

## Performance and Tuning Advantages

The two-kernel strategy yields **at least a 15% end-to-end speed-up** by eliminating idle SMs and improving occupancy across the device. Additionally, the separation allows each stage to be **independently tunable**—developers can optimize block sizes, shared-memory layouts, and occupancy strategies for K1 and K2 separately without sacrificing the other stage's performance. This provides a clean optimization surface that monolithic kernels cannot match.

## Implementation in the FlashKDA Codebase

The dispatch logic in [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp) orchestrates the two-kernel launch through the `launch_fwd` function. This function first executes `_flash_kda_fwd_prepare` (K1) and then `_flash_kda_fwd_recurrence` (K2), managing the handoff between the token-parallel and head-parallel stages.

```cpp
// Conceptual flow in csrc/flash_kda.cpp
launch_fwd() {
    _flash_kda_fwd_prepare(...);    // K1: Token-parallel grid N×H×num_chunks
    _flash_kda_fwd_recurrence(...); // K2: Head-parallel grid N×H
}

```

For Python integration, the high-level `chunk_kda` function from `fla.ops.kda` abstracts this two-kernel dispatch:

```python
import torch
from fla.ops.kda import chunk_kda

with torch.inference_mode():
    out, final_state = chunk_kda(
        q=q, k=k, v=v, g=g, beta=beta,
        scale=scale,
        initial_state=h0,
        output_final_state=True,
        use_gate_in_kernel=True,
        use_qk_l2norm_in_kernel=True,
        use_beta_sigmoid_in_kernel=True,
        safe_gate=True,
        A_log=A_log, dt_bias=dt_bias,
        lower_bound=lower_bound,
        transpose_state_layout=True,
        cu_seqlens=cu_seqlens,
    )

```

The low-level `fwd` API provides direct access to the two-kernel pipeline:

```python
import torch
from flash_kda import fwd

fwd(
    q, k, v, g, beta,
    scale,
    out,
    workspace,
    A_log, dt_bias,
    lower_bound,
    initial_state=None,
    final_state=None,
    cu_seqlens=None,
)

```

## Summary

- **FlashKDA's two-kernel fusion strategy** splits Kimi Delta Attention into K1 (token-parallel) and K2 (head-parallel) stages to maximize GPU utilization according to their distinct parallelism requirements.
- **K1** uses a grid of `N × H × num_chunks` to parallelize preparation work (gate activation, normalization, matrix construction), while **K2** uses `N × H` for recurrence computation (delta-rule, state accumulation).
- **Separating the kernels** eliminates the "lowest common denominator" bottleneck of monolithic designs, preventing the high-parallelism K1 work from being throttled by the lower-parallelism K2 recurrence.
- This architecture delivers **at least 15% end-to-end speedup** and enables **independent tuning** of block sizes, shared-memory layouts, and occupancy strategies for each kernel.

## Frequently Asked Questions

### What is the FlashKDA two-kernel fusion strategy?

The FlashKDA two-kernel fusion strategy is an architectural approach that decomposes the Kimi Delta Attention computation into two distinct CUDA kernels: K1 for token-parallel preparation and K2 for head-parallel recurrence. K1 handles gate activation, L2-normalisation, decay calculations, and matrix construction with grid dimensions `N × H × num_chunks`, while K2 manages the delta-rule recurrence and state accumulation with grid `N × H`. This separation allows each stage to utilize its optimal parallelism pattern without constraining the other.

### How does the two-kernel approach improve GPU utilization over a single fused kernel?

A single fused kernel forces all operations to adopt the lowest common denominator of parallelism, causing the high-parallelism K1 work to be throttled by the lower-parallelism K2 recurrence. By splitting into K1 and K2, FlashKDA allows K1 to saturate thousands of SMs with token-level work while K2 runs independently with a smaller grid sized to its specific needs. This eliminates idle SMs and improves overall device occupancy.

### Where are the K1 and K2 kernels implemented in the FlashKDA codebase?

Kernel 1 is implemented in `csrc/smxx/fwd_kernel1.cuh` and handles token-parallel operations including matrix inversion and gate activation. Kernel 2 resides in `csrc/smxx/fwd_kernel2.cuh` and manages head-parallel recurrence processing. The dispatch logic that launches K1 followed by K2 is located in [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp) within the `launch_fwd` function, which calls `_flash_kda_fwd_prepare` and then `_flash_kda_fwd_recurrence`.

### Can K1 and K2 be optimized independently in the FlashKDA architecture?

Yes, the two-kernel separation enables independent optimization of block sizes, shared-memory layouts, and occupancy strategies for each stage. Developers can tune K1 for maximum throughput on highly parallel token-level workloads while separately optimizing K2 for efficient recurrence processing. This independent tunability is impossible in a monolithic kernel design, where optimizations for one section inherently constrain the other.