Which PyTorch Version is Compatible with FlashKDA? Requirements and Setup Guide
FlashKDA requires PyTorch 2.4 or newer with CUDA support, as explicitly stated in the repository's README and enforced during the build process.
FlashKDA is a high-performance CUDA kernel library developed by MoonshotAI for accelerating deep learning workloads. To leverage these optimized kernels, your environment must meet specific PyTorch version constraints that ensure compatibility with the underlying CUDA extensions.
Understanding FlashKDA PyTorch Version Requirements
The MoonshotAI/FlashKDA repository maintains strict version requirements to guarantee stable operation across different hardware configurations.
Official Version Constraint
According to the README.md in the repository root, the Requirements section explicitly lists "PyTorch 2.4 and above" as a mandatory prerequisite【/tmp/instagit_8zrju49u/README.md†L9-L13】. This constraint ensures access to specific PyTorch APIs and CUDA features that FlashKDA's kernels depend upon.
CUDA Compilation Requirement
Beyond the version number, FlashKDA requires a PyTorch installation compiled with CUDA support. In setup.py, the build system uses torch.utils.cpp_extension.CUDAExtension to compile the custom kernels, which validates that your PyTorch distribution includes CUDA capabilities【/tmp/instagit_8zrju49u/setup.py†L33-L34】. CPU-only PyTorch installations will fail during the build phase.
Verifying Your PyTorch Installation
Before installing FlashKDA, confirm your PyTorch version and CUDA availability:
import torch
# Check PyTorch version meets the 2.4+ requirement
version = torch.__version__.split('+')[0] # Remove CUDA suffix if present
major, minor = map(int, version.split('.')[:2])
assert (major > 2) or (major == 2 and minor >= 4), "FlashKDA requires PyTorch 2.4 or newer"
# Verify CUDA is available
assert torch.cuda.is_available(), "FlashKDA requires PyTorch compiled with CUDA support"
print(f"PyTorch {torch.__version__} with CUDA {torch.version.cuda} is compatible with FlashKDA")
Installing FlashKDA with Compatible PyTorch
Ensure you install PyTorch 2.4 or newer with CUDA support before attempting to build FlashKDA:
- Install compatible PyTorch:
pip install torch>=2.4.0 --index-url https://download.pytorch.org/whl/cu121 - Clone the FlashKDA repository:
git clone https://github.com/MoonshotAI/FlashKDA.git - Install FlashKDA:
cd FlashKDA && pip install -e .
The setup.py build process will automatically detect whether your PyTorch installation includes the necessary CUDA headers and libraries required to compile the extension.
Code Example: Running FlashKDA with Validated PyTorch
Once you have PyTorch 2.4+ with CUDA support installed, you can import and execute the fwd function exposed in flash_kda/__init__.py:
import torch
from flash_kda import fwd
# Verify version compatibility at runtime
assert torch.__version__ >= "2.4", "FlashKDA requires PyTorch 2.4 or newer"
# Configure tensor dimensions
B, T, H, K = 1, 128, 8, 128
dtype = torch.bfloat16
device = torch.device("cuda")
# Initialize input tensors (FlashKDA operates on bf16/fp16 tensors)
q = torch.randn(B, T, H, K, dtype=torch.float32, device=device).to(dtype)
k = torch.randn(B, T, H, K, dtype=torch.float32, device=device).to(dtype)
v = torch.randn(B, T, H, K, dtype=torch.float32, device=device).to(dtype)
g = torch.randn(B, T, H, K, dtype=torch.float32, device=device).to(dtype)
beta = torch.randn(B, T, H, dtype=torch.float32, device=device).to(dtype)
scale = 1.0
out = torch.empty_like(q)
A_log = torch.rand(H, dtype=torch.float32, device=device)
dt_bias = torch.rand(H, K, dtype=torch.float32, device=device)
lower_bound = -5.0
# Execute the optimized forward kernel
fwd(q, k, v, g, beta, scale, out, A_log, dt_bias, lower_bound)
print("FlashKDA forward operation completed successfully")
This example demonstrates the tensor shapes and data types expected by the fwd function, which performs the accelerated kernel dispatch operation that FlashKDA provides.
Summary
- PyTorch 2.4 or newer is required to run FlashKDA, as documented in
README.md【/tmp/instagit_8zrju49u/README.md†L9-L13】. - CUDA support is mandatory—the
setup.pybuild system validates CUDA compilation viaCUDAExtension【/tmp/instagit_8zrju49u/setup.py†L33-L34】. - Validate your installation using
torch.__version__andtorch.cuda.is_available()before building. - The main entry point
fwdinflash_kda/__init__.pyaccepts bf16/fp16 tensors and requires a CUDA-enabled PyTorch environment.
Frequently Asked Questions
What is the minimum PyTorch version for FlashKDA?
FlashKDA requires PyTorch 2.4 as the absolute minimum version. This requirement is explicitly stated in the repository's README and ensures compatibility with the CUDA kernel interfaces that FlashKDA utilizes.
Does FlashKDA work with CPU-only PyTorch installations?
No. FlashKDA requires PyTorch compiled with CUDA support. The build process in setup.py uses torch.utils.cpp_extension.CUDAExtension, which will fail if your PyTorch installation lacks CUDA libraries, even if you meet the version 2.4+ requirement.
How do I check if my PyTorch installation is compatible with FlashKDA?
Run python -c "import torch; print(torch.__version__); print(torch.cuda.is_available())". You should see a version string of 2.4.0 or higher and True for CUDA availability. If either check fails, you must reinstall PyTorch with the appropriate CUDA-enabled wheel.
Will FlashKDA work with PyTorch 2.5 or newer versions?
Yes. FlashKDA is compatible with any PyTorch version greater than or equal to 2.4, including future releases like 2.5, 2.6, and beyond, provided they maintain the CUDA extension APIs that FlashKDA depends upon. The >= 2.4 constraint in the README indicates forward compatibility.
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 →