# How to Use FlashKDA as a Backend for flash-linear-attention: Complete Integration Guide

> Integrate FlashKDA as a high-performance CUDA backend for flash-linear-attention. Learn how to enable and leverage optimized Kimi Delta Attention kernels on supported GPUs.

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

---

**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:

```bash

# 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`](https://github.com/MoonshotAI/FlashKDA/blob/main/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`](https://github.com/MoonshotAI/FlashKDA/blob/main/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`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp) (lines 62–110) enforces strict dtype and memory layout constraints. The function signature is:

```python
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]` where `K = V = 128`
- **Beta**: `bfloat16` tensor of shape `[B, T, H]`
- **A_log**: `float32` tensor of shape `[H]`
- **dt_bias**: `float32` tensor of shape `[H, 128]`
- **lower_bound**: Python `float` in 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:

```python
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]` where `B` is the batch size
- Each sequence in the batch must be equal length

**Variable-Length Mode** (`cu_seqlens` provided):
- State shape: `[N, H, V, K]` where `N` is the total number of sequences
- Input batch dimension `B` must equal `1`
- `cu_seqlens` is an integer tensor marking sequence boundaries

The input validation logic in [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/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`](https://github.com/MoonshotAI/FlashKDA/blob/main/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):

```bash
export FLA_FLASH_KDA=0

```

**To ensure FlashKDA is selected** (explicit opt-in):

```bash
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_KDA` environment variable controls dispatch; it defaults to enabled (1)
- FlashKDA requires `bfloat16` inputs with head dimension 128 and `float32` for `A_log`/`dt_bias`
- Use `chunk_kda` from `fla.ops.kda` with `torch.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=0` to 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`](https://github.com/MoonshotAI/FlashKDA/blob/main/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.