# How FlashKDA Ensures Portability Across Modern NVIDIA GPUs Using SM80 MMA Instructions

> Discover how FlashKDA ensures portability across NVIDIA GPUs. This article details its use of SM80 MMA instructions and CUTLASS abstraction for efficient PTX generation on Ampere Hopper and Blackwell.

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

---

**FlashKDA achieves cross-architecture portability by layering CUTLASS abstraction templates atop SM80 MMA primitives, enabling automatic generation of optimal PTX instructions for Ampere, Hopper, and Blackwell GPUs at compile time.**

FlashKDA is an open-source CUDA kernel library developed by MoonshotAI that accelerates attention mechanisms using NVIDIA’s Tensor Core MMA instructions. The repository guarantees **SM80 MMA portability** across modern hardware by decoupling algorithm logic from architecture-specific instruction selection, allowing a single C++ template codebase to target GPUs from SM80 through SM10x without source modification.

## CUTLASS Foundation for MMA Abstraction

FlashKDA builds directly on NVIDIA’s CUTLASS library to abstract low-level matrix-multiply-accumulate primitives. Rather than hand-writing PTX assembly for each GPU generation, the codebase delegates instruction selection to CUTLASS architecture tags.

In `csrc/smxx/fwd_kernel1.cuh` and `csrc/smxx/fwd_kernel2.cuh`, kernels forward compile-time architecture parameters to CUTLASS wrapper types:

- **`cutlass::arch::Sm80`** maps to Ampere’s `wmma` PTX instructions
- **`cutlass::arch::Sm90`** generates Hopper’s `mma` async operations  
- **`cutlass::arch::OpClassTensorOp`** automatically selects the appropriate Tensor Core width

This delegation ensures that when the compiler targets SM80, it emits classic `wmma.mma.sync.aligned` instructions, while SM90 targets receive `mma.sync.aligned` with support for floating-point accumulate types—all from identical source templates.

## Compile-Time Architecture Selection

The portability strategy relies on C++ template metaprogramming to eliminate runtime overhead. FlashKDA kernels use architecture-agnostic templating where the concrete SM version is resolved during compilation.

### Template Parameter Propagation

All kernel implementations in `csrc/smxx/fwd_kernel1.cuh` declare architecture-dependent types via template parameters:

```cpp
template <typename Arch>
struct KernelTraits {
    using MmaShape = typename cutlass::gemm::GemmShape<...>;
    using ArchTag = Arch;
};

```

The `Arch` parameter is instantiated as `cutlass::arch::Sm80`, `cutlass::arch::Sm90`, or future SM10x tags based on the build configuration. This allows the same kernel body to leverage SM90-specific features like TMA (Tensor Memory Accelerator) when available, while maintaining backward compatibility with SM80’s narrower register files.

### Feature Selection via Preprocessor

Inside kernel files, optimal tile sizes and warp layouts are selected using compile-time conditionals:

```cpp
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
    // Hopper-optimized tile configuration
    constexpr int kTileSizeM = 256;
#else
    // Ampere-compatible configuration  
    constexpr int kTileSizeM = 128;
#endif

```

This ensures that each GPU receives the most efficient MMA tile configuration without requiring separate source files for each architecture.

## Build System and Runtime Detection

The [`setup.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/setup.py) build script implements automatic device detection to streamline deployment across heterogeneous environments.

### Environment-Driven Compilation

Users control target architectures via the `FLASH_KDA_CUDA_ARCHS` environment variable:

```bash

# Compile for specific architecture

FLASH_KDA_CUDA_ARCHS=80,90 pip install --no-build-isolation .

# Build universal wheel containing all supported SM versions

FLASH_KDA_CUDA_ARCHS=all pip install --no-build-isolation .

```

When set to `all`, the build system injects multiple `-gencode` flags (e.g., `arch=compute_80,code=sm_80`, `arch=compute_90,code=sm_90`) into the NVCC command line, generating distinct cubin objects for each architecture within a single binary.

### Host-Side Dispatch Logic

The `csrc/smxx/fwd_launch.cu` file implements the host-side launch wrapper that queries the CUDA driver for the current device’s compute capability. At runtime, the dispatcher selects the optimal kernel specialization based on the detected SM version, ensuring that an A100 (SM80) executes the Ampere-optimized path while an H100 (SM90) utilizes Hopper-specific instructions without user intervention.

## Graceful Degradation and Fallback

FlashKDA incorporates defensive programming to handle hardware capability mismatches. If a target GPU does not support the requested MMA width or instruction format, CUTLASS dispatches to a standard CUDA arithmetic implementation. This fallback guarantees functional correctness across the entire SM80+ family, though with reduced performance compared to Tensor Core paths.

The Python binding layer in [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp) exposes these kernels through a unified `flash_kda.fwd` interface that hides architecture complexity from PyTorch users:

```python
import torch
from flash_kda import fwd

# Data preparation (bf16 tensors)

B, T, H, K, V = 2, 128, 8, 128, 128
q = torch.randn(B, T, H, K, dtype=torch.bfloat16, device='cuda')
k = torch.randn_like(q)
v = torch.randn(B, T, H, V, dtype=torch.bfloat16, device='cuda')
g = torch.randn_like(q)
beta = torch.randn(B, T, H, dtype=torch.bfloat16, device='cuda')

# Output buffers

out = torch.empty_like(v)
A_log = torch.zeros(H, dtype=torch.float32, device='cuda')
dt_bias = torch.zeros(H, K, dtype=torch.float32, device='cuda')

# Architecture-agnostic kernel invocation

flash_kda.fwd(
    q, k, v, g, beta, scale=0.125,
    out=out, A_log=A_log, dt_bias=dt_bias, 
    lower_bound=-5.0,
    initial_state=None, final_state=None, cu_seqlens=None
)

```

## Summary

- **CUTLASS Integration**: FlashKDA uses `cutlass::arch::Sm80` and `cutlass::arch::Sm90` tags to abstract MMA instruction generation, automatically emitting correct PTX for each architecture.
- **Template-Based Polymorphism**: Architecture selection occurs at compile time via C++ templates in `fwd_kernel1.cuh` and `fwd_kernel2.cuh`, eliminating runtime branching overhead.
- **Multi-Architecture Binaries**: The [`setup.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/setup.py) build system supports `FLASH_KDA_CUDA_ARCHS=all` to generate fat binaries containing optimized cubins for SM80, SM90, and SM10x devices.
- **Runtime Dispatch**: `fwd_launch.cu` detects the host GPU compute capability and selects the corresponding kernel specialization automatically.
- **Fallback Safety**: CUTLASS provides automatic fallback to standard CUDA arithmetic when Tensor Core operations are unavailable, ensuring functional portability.

## Frequently Asked Questions

### What specific CUDA architectures does FlashKDA support?

FlashKDA officially supports **SM80 (Ampere)** and newer, including **SM90 (Hopper)** and **SM10x (Blackwell)** architectures. The codebase is structured to accept future architecture tags through the CUTLASS abstraction layer without requiring kernel rewrites.

### How does FlashKDA handle different MMA instruction formats between GPU generations?

The library delegates instruction format selection to CUTLASS architecture wrappers. When compiled for SM80, CUTLASS generates `wmma` PTX instructions; for SM90, it emits `mma` operations with support for asynchronous pipelines. This mapping occurs transparently during template instantiation in `fwd_kernel1.cuh`.

### Can I compile FlashKDA for multiple GPU architectures simultaneously?

Yes. Set the environment variable `FLASH_KDA_CUDA_ARCHS=all` during installation. This instructs the [`setup.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/setup.py) build script to pass multiple `-gencode` flags to NVCC, embedding separate cubin objects for SM80, SM90, and other supported architectures within the wheel. The CUDA runtime automatically selects the optimal binary at kernel launch.

### What happens if I run FlashKDA on a GPU older than SM80?

FlashKDA requires SM80 or newer due to its reliance on Tensor Core MMA instructions. On unsupported hardware, the build system will fail compilation unless `FLASH_KDA_CUDA_ARCHS` explicitly targets the older architecture, in which case CUTLASS falls back to SIMT (single-instruction, multiple-thread) implementations with significantly reduced performance.