How to Specify Custom CUDA Architectures When Compiling FlashKDA

Set the FLASH_KDA_CUDA_ARCHS environment variable to auto, all, or a comma-separated list of compute capabilities (e.g., 90a,100a) before running pip install.

FlashKDA accelerates attention mechanisms with custom CUDA kernels that must be compiled for specific GPU architectures. The build system located in setup.py supports flexible architecture targeting through environment variables, allowing you to optimize for a single GPU or maintain compatibility across multiple hardware generations.

How Architecture Selection Works in setup.py

The compilation logic resides in setup.py and centers on two key components: the SUPPORTED_CUDA_ARCHS list and the get_arch_flags() helper function.

At lines 19–20, FlashKDA defines its supported architectures:

SUPPORTED_CUDA_ARCHS = ["90a", "100a", "103a", "120a"]

Between lines 35–47, the build script reads the FLASH_KDA_CUDA_ARCHS environment variable. It supports three resolution modes:

  • auto (default): Detects the capability of the currently visible GPU using torch.cuda.get_device_capability
  • all: Selects every architecture listed in SUPPORTED_CUDA_ARCHS
  • Custom list: Parses a comma-separated string provided by the user

Lines 49–52 convert each architecture string into NVCC -gencode flags. For example, 90a becomes:

-gencode arch=compute_90a,code=sm_90a

Finally, at lines 68–84, these flags populate extra_compile_args['nvcc'] and pass into torch.utils.cpp_extension.CUDAExtension.

Setting Custom Architectures via Environment Variables

To specify custom CUDA architectures when compiling FlashKDA, export FLASH_KDA_CUDA_ARCHS before invoking the installer. The variable accepts three distinct formats.

Automatic Detection (Default)

When no variable is set, or set to auto, the build targets only the GPU currently attached to your system:

pip install flash_kda

Build for All Supported Architectures

Use all to generate fat binaries compatible with every supported generation. This creates four separate -gencode entries (sm_90a, sm_100a, sm_103a, sm_120a):

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

Target Specific Architectures

Supply a comma-separated list of compute capabilities using the a suffix (indicating Ampere-compatible feature sets):


# Compile for SM 80 (CUDA 11.x) and SM 90 (CUDA 12.x)

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

You can combine this with other PyTorch extension variables like NVCC_THREADS:

NVCC_THREADS=8 FLASH_KDA_CUDA_ARCHS=90a pip install -v --no-build-isolation .

Understanding Architecture Codes

FlashKDA uses architecture strings with an a suffix to denote features compatible with Ampere and newer GPUs. The following codes are recognized in SUPPORTED_CUDA_ARCHS:

  • 90a: Compute capability 9.0 (SM 90) — RTX 4090, H100
  • 100a: Compute capability 10.0 (SM 100) — Hopper A100
  • 103a: Compute capability 10.3 (SM 103) — Future Hopper-based GPUs
  • 120a: Compute capability 12.0 (SM 120) — Next-generation architectures (placeholder)

If you supply an architecture string absent from SUPPORTED_CUDA_ARCHS, NVCC will raise an error during compilation. Ensure your target hardware supports the requested compute capability.

Optional: Clang Tooling Configuration

While setup.py controls the actual compilation, config.yaml (lines 14–20) provides default flags for clang-based language servers and IDEs. This file defaults to sm_90, but the build process always defers to the NVCC flags generated by get_arch_flags() during actual package installation.

Summary

  • Primary control: Set FLASH_KDA_CUDA_ARCHS to auto, all, or a comma-separated list (e.g., 90a,100a).
  • Source location: Architecture logic resides in setup.py (lines 19–84), specifically within SUPPORTED_CUDA_ARCHS and get_arch_flags().
  • Flag generation: Valid entries convert to -gencode arch=compute_XXa,code=sm_XXa flags passed to NVCC via extra_compile_args.
  • Supported values: 90a, 100a, 103a, and 120a cover modern NVIDIA GPUs from Ampere through future Hopper generations.
  • Validation: Unsupported architecture strings cause NVCC to fail at compile time.

Frequently Asked Questions

What happens if I do not set FLASH_KDA_CUDA_ARCHS?

The build system defaults to auto mode. It queries torch.cuda.get_device_capability to detect the compute capability of the currently visible GPU and compiles exclusively for that architecture. This produces lean binaries but may fail if you later run the package on different hardware.

Can I compile FlashKDA for older architectures like SM 7.0 or SM 8.0?

The SUPPORTED_CUDA_ARCHS list in setup.py officially recognizes 90a, 100a, 103a, and 120a. While you might attempt to pass 80a or similar values, the build script validates entries against this list. For unsupported legacy hardware, you would need to modify the source code or use pre-built wheels if available.

Why do the architecture codes include an "a" suffix?

The "a" suffix denotes Ampere-compatible feature sets and instruction sets required by FlashKDA's optimized kernels. When converted to NVCC flags, 90a expands to -gencode arch=compute_90a,code=sm_90a, ensuring the compiler targets the correct microarchitecture variant.

How can I verify which architectures were compiled into my FlashKDA installation?

After installation, inspect the build logs generated by pip install -v. The verbose output lists the exact -gencode flags passed to NVCC during compilation. Alternatively, run torch.utils.cpp_extension.load with verbose logging to confirm the resulting binary contains code for your target SM versions.

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 →