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

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:

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:

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


# 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 exposes these kernels through a unified flash_kda.fwd interface that hides architecture complexity from PyTorch users:

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 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 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.

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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →