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:

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

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 →