How FlashKDA Integrates with CUTLASS for High-Performance GEMM Operations
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. 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:
-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, 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:
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, lines 20–25 convert PyTorch at::BFloat16 pointers to CUTLASS’s native bfloat16_t type using reinterpret_cast:
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 (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:
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.
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:
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:
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:
# 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.yamlenables SM90 Tensor-Core support via-DCUTLASS_ARCH_MMA_SM90_SUPPORTED=1for CUTLASS MMA primitives. - Type Safety:
csrc/flash_kda.cppenforces BFloat16 inputs and converts pointers tocutlass::bfloat16_tfor zero-copy GEMM execution. - Template Dispatch: The
LAUNCHmacro routes tolaunch_fwdincsrc/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_sizeensures 128-byte aligned buffers required by CUTLASS tensor-core operations. - Variable Lengths:
cu_seqlenstriggers 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 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 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.
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 →