# How FlashKDA Integrates with CUTLASS for High-Performance GEMM Operations

> Discover how FlashKDA integrates with CUTLASS using a C++ interface and PyTorch BFloat16 tensors to dispatch templated CUDA kernels for high-performance GEMM operations.

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

---

**FlashKDA leverages NVIDIA’s CUTLASS library by wrapping its Tensor-Core GEMM primitives in a thin C++ interface that converts PyTorch BFloat16 tensors to CUTLASS types and dispatches to templated CUDA kernels.**

FlashKDA, developed by MoonshotAI, implements its K-Delta-Attention (KDA) algorithm using CUTLASS’s highly tuned matrix-multiply-accumulate (GEMM) building blocks. The integration occurs across the build system, tensor validation layer, and kernel launch pipeline, enabling efficient execution on SM90 (Hopper) architectures.

## Compile-Time Configuration and SM90 Support

FlashKDA enables CUTLASS Tensor-Core support through CMake configuration flags defined in [`config.yaml`](https://github.com/MoonshotAI/FlashKDA/blob/main/config.yaml). The build system explicitly targets SM90 (Hopper) architecture capabilities required for the underlying GEMM operations.

The configuration passes the following compiler flag twice to enable the necessary MMA (matrix-multiply accumulate) instructions:

```bash
-DCUTLASS_ARCH_MMA_SM90_SUPPORTED=1

```

This flag ensures that CUTLASS headers can instantiate templates utilizing the SM90-specific tensor core instructions that power FlashKDA’s attention mechanism.

## Data Type Constraints and Tensor Preparation

All input tensors must conform to strict BFloat16 requirements to align with CUTLASS’s native data layouts. In [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp), lines 49–55 enforce that query (`q`), key (`k`), value (`v`), gating (`g`), beta, and output tensors are CUDA BFloat16 tensors.

The validation logic checks:

```cpp
TORCH_CHECK(q.scalar_type() == torch::kBFloat16, "Input tensors must be BFloat16");

```

This constraint exists because the CUTLASS GEMM wrappers in `csrc/smxx/fwd_launch.cu` are instantiated specifically for `cutlass::bfloat16_t` types, matching the hardware-native tensor core precision.

## Pointer Conversion and CUTLASS Type Mapping

FlashKDA bridges PyTorch’s ATen tensors and CUTLASS’s type system through explicit pointer casting. In [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp), lines 20–25 convert PyTorch `at::BFloat16` pointers to CUTLASS’s native `bfloat16_t` type using `reinterpret_cast`:

```cpp
reinterpret_cast<cutlass::bfloat16_t const*>(q_3d.data_ptr<at::BFloat16>())

```

This conversion allows the CUTLASS device kernels to access tensor data directly without intermediate copies, maintaining zero-overhead data transfer between the PyTorch runtime and the GEMM implementations.

## Kernel Launch Pipeline and Template Dispatching

The forward pass (`fwd`) dispatches to specialized CUDA kernels through a templated launcher mechanism defined in `csrc/smxx/fwd_launch.cu`. The entry point in [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp) (lines 84–90) invokes the `LAUNCH` macro, which selects the appropriate `launch_fwd` instantiation based on head dimensions and precision requirements.

The launcher signature follows this pattern:

```cpp
launch_fwd<128, HI, HO, FP32, VL>()

```

Where:
- **128** represents the fixed head dimension
- **HI** and **HO** are input/output head dimensions
- **FP32** controls the accumulator precision
- **VL** enables variable-length sequence support

This template ultimately invokes CUTLASS device GEMM wrappers (e.g., `cutlass::gemm::device::Gemm`) to perform the matrix multiplications driving the KDA recurrence.

## State Management and Accumulator Types

FlashKDA supports optional state tensors for recurrent attention mechanisms. The `initial_state` and `final_state` tensors may use either BFloat16 or FP32 precision, controlled by the `state_fp32` flag detected around lines 59–74 in [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp).

When `FP32` is true, the launcher selects `cutlass::half_t` versus `float` accumulator types, affecting the internal GEMM accumulation precision. This flag propagates through the dispatch logic to `launch_fwd`, ensuring that CUTLASS uses the appropriate floating-point pipeline for state-heavy computations.

## Variable-Length Sequence Support

For non-uniform sequence lengths, FlashKDA adjusts its CUTLASS tiling strategy through the variable-length (`VL`) template parameter. When `cu_seqlens` is provided, the detection logic (lines 45–60) sets `is_varlen = true`, triggering the `DISPATCH_STATE(true)` path.

This adaptation allows the CUTLASS-based kernels to handle irregular memory access patterns efficiently by adjusting tile coordinates and warp-level synchronization based on the prefix-sum offsets provided in `cu_seqlens`.

## Workspace Allocation and Memory Alignment

Before invoking CUTLASS kernels, FlashKDA allocates intermediate buffers through `get_workspace_size` (lines 5–26). This function computes the required buffer size for intermediate GEMM results and prefix-sum operations, ensuring **128-byte alignment** compatible with CUTLASS’s tensor-core memory access patterns.

The workspace tensor must be passed as a `torch.uint8` buffer sized according to the batch and sequence length:

```python
workspace = torch.empty(flash_kda.get_workspace_size(B*T, H), dtype=torch.uint8, device='cuda')

```

## Practical Implementation Example

Below is a minimal Python example exercising the CUTLASS-backed forward kernel:

```python
import torch
import flash_kda

B, T, H, D = 2, 256, 16, 128               # B-batch, T-tokens, H-heads, D-dim

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')
workspace = torch.empty(flash_kda.get_workspace_size(B*T, H), dtype=torch.uint8, device='cuda')

out = torch.empty_like(q)

# Run the forward pass (CUTLASS GEMM kernels are invoked under the hood)

flash_kda.fwd(
    q, k, v, g, beta,
    scale=1.0,
    out=out,
    workspace=workspace,
    A_log=A_log,
    dt_bias=dt_bias,
    lower_bound=0.1
)

print(out.shape)   # → torch.Size([2, 256, 16, 128])

```

For stateful or variable-length execution:

```python

# Optional state (FP32 accumulator)

init_state = torch.randn(B, H, D, D, dtype=torch.float32, device='cuda')
final_state = torch.empty_like(init_state)

# Variable-length example

cu_seqlens = torch.tensor([0, 128, 256], dtype=torch.int64, device='cuda')

flash_kda.fwd(
    q, k, v, g, beta,
    scale=1.0,
    out=out,
    workspace=workspace,
    A_log=A_log,
    dt_bias=dt_bias,
    lower_bound=0.1,
    initial_state=init_state,
    final_state=final_state,
    cu_seqlens=cu_seqlens
)

```

## Summary

- **Build Configuration**: [`config.yaml`](https://github.com/MoonshotAI/FlashKDA/blob/main/config.yaml) enables SM90 Tensor-Core support via `-DCUTLASS_ARCH_MMA_SM90_SUPPORTED=1` for CUTLASS MMA primitives.
- **Type Safety**: [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp) enforces BFloat16 inputs and converts pointers to `cutlass::bfloat16_t` for zero-copy GEMM execution.
- **Template Dispatch**: The `LAUNCH` macro routes to `launch_fwd` in `csrc/smxx/fwd_launch.cu`, which instantiates CUTLASS device GEMM templates.
- **Precision Control**: State tensors support both BFloat16 and FP32 accumulators, toggling CUTLASS accumulator types via template parameters.
- **Memory Management**: `get_workspace_size` ensures 128-byte aligned buffers required by CUTLASS tensor-core operations.
- **Variable Lengths**: `cu_seqlens` triggers variable-length mode, adapting CUTLASS tiling strategies for non-uniform sequences.

## Frequently Asked Questions

### What specific CUTLASS components does FlashKDA use for its GEMM operations?

FlashKDA utilizes `cutlass::gemm::device::Gemm` classes and MMA (matrix-multiply accumulate) operators specifically tuned for SM90 architectures. The integration is visible in `csrc/smxx/fwd_launch.cu` and `csrc/smxx/utils.cuh`, where device-side helper functions wrap CUTLASS MMA instructions to perform the core matrix multiplications within the K-Delta-Attention recurrence.

### Why does FlashKDA require BFloat16 tensors for CUTLASS integration?

CUTLASS’s SM90 Tensor-Core GEMM primitives are optimized for `bfloat16_t` arithmetic, which provides the optimal balance of numerical range and computational throughput on Hopper GPUs. The `torch::kBFloat16` checks in [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp) ensure data alignment with CUTLASS’s native types before pointer casting, preventing type mismatches during kernel execution.

### How does FlashKDA handle variable-length sequences with CUTLASS?

When a `cu_seqlens` tensor is provided, FlashKDA detects variable-length mode (`is_varlen`) and sets the `VL` template parameter to true in the `launch_fwd` dispatcher. This modifies the CUTLASS kernel’s tiling and indexing logic to accommodate irregular sequence boundaries while maintaining coalesced memory access patterns through the prefix-sum offsets.

### Can FlashKDA use FP32 accumulation with CUTLASS GEMMs?

Yes. FlashKDA supports FP32 accumulator precision for state tensors (`initial_state`, `final_state`). The `state_fp32` boolean flag in [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp) propagates to the `launch_fwd` template, selecting `float` versus `cutlass::half_t` as the CUTLASS accumulator type. This allows higher precision for recurrent state updates while keeping input/output tensors in BFloat16.