How to Use FlashKDA as a Backend for flash-linear-attention: Complete Integration Guide
FlashKDA automatically serves as the high-performance CUDA backend for flash-linear-attention when the FLA_FLASH_KDA environment variable is enabled (the default), routing chunk_kda calls to optimized Kimi Delta Attention kernels on supported GPUs.
MoonshotAI's FlashKDA repository provides a production-grade CUDA implementation of Kimi Delta Attention (KDA) that integrates seamlessly with the flash-linear-attention (FLA) library. When properly configured, FLA's Python API automatically dispatches attention computations to FlashKDA's fused kernels rather than Triton fallbacks, delivering substantial throughput improvements on modern NVIDIA hardware. This guide covers the installation process, dispatch mechanism, and correct usage patterns for inference workloads.
Installation and Setup
Using FlashKDA as a backend requires installing both the FLA library and the FlashKDA CUDA extension. FlashKDA must be built from source to ensure compatibility with your local CUDA toolkit.
Install the dependencies in the following order:
# Install flash-linear-attention (v0.5.0 or newer)
pip install -U flash-linear-attention
# Build and install FlashKDA from the repository root
pip install -v --no-build-isolation .
The setup.py file in the FlashKDA repository handles the compilation of the flash_kda_C extension with appropriate architecture flags for SM 90+ GPUs.
Understanding the Dispatch Mechanism
The integration relies on an environment variable check within FLA's chunk_kda wrapper. According to the source code in flash_kda/__init__.py (lines 5–42), the wrapper inspects FLA_FLASH_KDA at runtime to determine which implementation to use.
If FLA_FLASH_KDA is set to 1 or unset (default), the call routes to flash_kda.fwd. If set to 0, the wrapper falls back to the Triton implementation.
The dispatch logic passes identical tensor arguments to both backends, including optional state tensors and cu_seqlens for variable-length batching. When successful, the FLA logger outputs:
[FLA Backend] kda.chunk_kda -> flashkda
Kernel API and Tensor Requirements
The underlying CUDA kernel exposed via PyBind11 in csrc/flash_kda.cpp (lines 62–110) enforces strict dtype and memory layout constraints. The function signature is:
flash_kda.fwd(q, k, v, g, beta, scale, out,
A_log, dt_bias, lower_bound,
initial_state=None, final_state=None, cu_seqlens=None)
Input specifications:
- Main tensors (
q,k,v,g): Must be CUDA-contiguous,bfloat16, and shape[B, T, H, 128]whereK = V = 128 - Beta:
bfloat16tensor of shape[B, T, H] - A_log:
float32tensor of shape[H] - dt_bias:
float32tensor of shape[H, 128] - lower_bound: Python
floatin the range[-5.0, 0]
The kernel validates these shapes and dtypes internally before launching the CUDA stream.
Running Inference with chunk_kda
For production inference, wrap calls in torch.inference_mode() to disable gradient computation. The following example demonstrates the standard batched mode:
import torch
import logging
from fla.ops.kda import chunk_kda
# Enable backend logging to verify dispatch
logging.basicConfig(level=logging.INFO)
# Configuration
B, T, H, D = 2, 256, 8, 128
scale = 0.125
lower_bound = -3.0
# Allocate dummy tensors (replace with real data)
q = torch.randn(B, T, H, D, dtype=torch.bfloat16, device='cuda')
k = torch.randn_like(q)
v = torch.randn_like(q)
g = torch.randn_like(q)
beta = torch.randn(B, T, H, dtype=torch.bfloat16, device='cuda')
A_log = torch.randn(H, dtype=torch.float32, device='cuda')
dt_bias = torch.randn(H, D, dtype=torch.float32, device='cuda')
# Optional initial state for recurrent processing
initial_state = torch.zeros(B, H, D, D, dtype=torch.bfloat16, device='cuda')
with torch.inference_mode():
out, final_state = chunk_kda(
q=q, k=k, v=v, g=g, beta=beta,
scale=scale,
initial_state=initial_state,
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=None, # Batched mode; omit for variable-length
)
print(f"Output shape: {out.shape}") # [B, T, H, D]
print(f"Final state shape: {final_state.shape}") # [B, H, D, D]
State Handling and Variable-Length Batches
FlashKDA supports both standard batched processing and variable-length sequences via the cu_seqlens parameter. The state tensor shape changes depending on the mode:
Batched Mode (cu_seqlens=None):
- State shape:
[B, H, V, K]whereBis the batch size - Each sequence in the batch must be equal length
Variable-Length Mode (cu_seqlens provided):
- State shape:
[N, H, V, K]whereNis the total number of sequences - Input batch dimension
Bmust equal1 cu_seqlensis an integer tensor marking sequence boundaries
The input validation logic in csrc/flash_kda.cpp (lines 62–110) ensures these constraints are met before kernel launch.
Performance Characteristics and Hardware Requirements
FlashKDA is optimized specifically for NVIDIA GPUs with Compute Capability 9.0 (SM 90) and newer, such as the H100 and H20 series. On supported hardware, the fused CUDA kernels significantly outperform the Triton reference implementation, particularly for long sequences and large batch sizes.
Benchmark results comparing throughput against the Triton fallback are available in the repository's BENCHMARK_H20.md file. For older GPUs (SM 80 and below), the Triton backend may provide better compatibility despite lower peak performance.
Debugging and Backend Verification
To verify that FlashKDA is active, enable Python logging at the INFO level as shown in the inference example. Look for the dispatch confirmation message in stderr.
To force the Triton implementation (for debugging or compatibility testing):
export FLA_FLASH_KDA=0
To ensure FlashKDA is selected (explicit opt-in):
export FLA_FLASH_KDA=1
If the CUDA extension fails to import or architecture mismatches occur, the wrapper automatically falls back to Triton regardless of the environment variable setting.
Summary
- Install both
flash-linear-attention(≥0.5.0) and FlashKDA from source to enable the backend - The
FLA_FLASH_KDAenvironment variable controls dispatch; it defaults to enabled (1) - FlashKDA requires
bfloat16inputs with head dimension 128 andfloat32forA_log/dt_bias - Use
chunk_kdafromfla.ops.kdawithtorch.inference_mode()for optimal inference performance - State tensors use shape
[B, H, V, K]for batched mode or[N, H, V, K]for variable-length sequences - Optimal performance requires SM 90+ GPUs (Hopper architecture); set
FLA_FLASH_KDA=0to force Triton on older hardware
Frequently Asked Questions
What hardware is required to run FlashKDA?
FlashKDA requires NVIDIA GPUs with Compute Capability 9.0 (SM 90) or newer, such as the H100 or H20. The kernels are specifically optimized for the Hopper architecture and will not compile or run on older GPUs like the A100 (SM 80). For unsupported hardware, use FLA_FLASH_KDA=0 to force the Triton fallback.
How can I confirm that FlashKDA is being used instead of the Triton backend?
Enable INFO-level logging by calling logging.basicConfig(level=logging.INFO) before your inference call. When FlashKDA is successfully dispatched, you will see the log message [FLA Backend] kda.chunk_kda -> flashkda. If this message does not appear, the system is using the Triton implementation.
Does FlashKDA support variable-length sequences in a batch?
Yes. FlashKDA supports variable-length sequences when you provide the cu_seqlens tensor, which contains the cumulative sequence lengths. In this mode, the batch dimension of input tensors must be 1, and the state tensor shape becomes [N, H, V, K] where N is the number of sequences. The standard batched mode (fixed-length sequences) uses cu_seqlens=None and state shape [B, H, V, K].
Why am I getting dtype errors when calling chunk_kda?
FlashKDA enforces strict type constraints validated in csrc/flash_kda.cpp. Ensure that query, key, value, and gate tensors are torch.bfloat16 and contiguous, while A_log and dt_bias must be torch.float32. Additionally, the head dimension must be exactly 128, as this is the only size supported by the current CUDA kernels.
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 →