# Which PyTorch Version is Compatible with FlashKDA? Requirements and Setup Guide

> Discover the compatible PyTorch version for FlashKDA. Learn the essential requirements and setup steps to use FlashKDA with PyTorch 2.4 and CUDA support.

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

---

**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`](https://github.com/MoonshotAI/FlashKDA/blob/main/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`](https://github.com/MoonshotAI/FlashKDA/blob/main/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:

```python
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:

1. Install compatible PyTorch: `pip install torch>=2.4.0 --index-url https://download.pytorch.org/whl/cu121`
2. Clone the FlashKDA repository: `git clone https://github.com/MoonshotAI/FlashKDA.git`
3. Install FlashKDA: `cd FlashKDA && pip install -e .`

The [`setup.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/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`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py):

```python
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`](https://github.com/MoonshotAI/FlashKDA/blob/main/README.md)【/tmp/instagit_8zrju49u/README.md†L9-L13】.
- **CUDA support is mandatory**—the [`setup.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/setup.py) build system validates CUDA compilation via `CUDAExtension`【/tmp/instagit_8zrju49u/setup.py†L33-L34】.
- Validate your installation using `torch.__version__` and `torch.cuda.is_available()` before building.
- The main entry point `fwd` in [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py) accepts 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`](https://github.com/MoonshotAI/FlashKDA/blob/main/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.