FlashKDA Two-Kernel Fusion Strategy: Why Splitting K1 and K2 Beats a Single CUDA Kernel
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 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 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.
// 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:
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:
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_chunksto parallelize preparation work (gate activation, normalization, matrix construction), while K2 usesN × Hfor 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 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.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →