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 usingtorch.cuda.get_device_capabilityall: Selects every architecture listed inSUPPORTED_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, H100100a: Compute capability 10.0 (SM 100) — Hopper A100103a: Compute capability 10.3 (SM 103) — Future Hopper-based GPUs120a: 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_ARCHStoauto,all, or a comma-separated list (e.g.,90a,100a). - Source location: Architecture logic resides in
setup.py(lines 19–84), specifically withinSUPPORTED_CUDA_ARCHSandget_arch_flags(). - Flag generation: Valid entries convert to
-gencode arch=compute_XXa,code=sm_XXaflags passed to NVCC viaextra_compile_args. - Supported values:
90a,100a,103a, and120acover 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →