# FlashKDA Correctness Tests: Validation Suite and Execution Guide

> Explore FlashKDA correctness tests. Learn how our pytest suite validates CUDA kernels against Python references, covering various data types and sequence lengths up to 1M tokens. Access the MoonshotAI/FlashKDA repo for executio...

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

---

**FlashKDA validates its custom CUDA kernel through exhaustive pytest suites that assert exact tensor equality against a double-precision Python reference implementation, covering state I/O variants, data types, sequence lengths up to 1M tokens, and variable-length batches.**

FlashKDA is a high-performance CUDA implementation of the Kernel Density Attention mechanism developed by MoonshotAI. To guarantee numerical accuracy across diverse deployment scenarios, the repository ships a comprehensive test suite that rigorously compares every kernel output against a pure-PyTorch reference running in FP64 double precision.

## What the FlashKDA Correctness Tests Cover

The validation strategy spans from minimal smoke tests to exhaustive parametrized matrices, ensuring correctness across state management modes, tensor layouts, and sequence lengths.

### Minimal Sanity Verification

The file [`tests/test_fwd.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/test_fwd.py) provides a deterministic smoke test for rapid iteration. It initializes a fixed random seed, constructs input tensors, and executes both `flash_kda.fwd` (the CUDA kernel exposed in [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py)) and the reference `torch_ref` function. The test asserts `torch.equal` on both the main output tensor and the final hidden state, demanding exact bitwise equality between the kernel and the reference on a single GPU with fixed sequence lengths.

### Exhaustive Parametrized Matrix

The file [`tests/test_fwd_full.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/test_fwd_full.py) contains the comprehensive correctness validation using `pytest.mark.parametrize` to sweep across critical dimensions:

- **State I/O Combinations**: Four distinct modes governing initial and final state tensors—`in+out` (both provided and returned), `in_only` (consume initial state only), `out_only` (produce final state only), and `no_state` (stateless operation).
- **State Data Types**: Both `bf16` (bfloat16) and `fp32` (float32) precision for state tensors.
- **Head Count (`H`)**: Representative values including 1, 4, 32, and 96 heads.
- **Sequence Length (`T`)**: Short sequences (≤1024), odd-length sequences, and extreme long-context scenarios up to 1,048,576 tokens.
- **Variable-Length Batches**: Ragged sequence handling via `cu_seqlens` tensors, including very long variable-length inputs.
- **Batch Dimension (`B`)**: Multi-sample batches with B ∈ {2, 4, 8} to verify automatic cu-seqlens handling.

Each configuration runs the kernel implemented in [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp) against the Python reference, asserting exact equality on both the output activation and the final state (when applicable).

### Edge Cases and Batching

Specific test functions within [`test_fwd_full.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/test_fwd_full.py) target boundary conditions:
- `test_fwd_varlen` exercises ragged batching logic.
- `test_fwd_long_varlen` validates sequences approaching the 1M token limit with variable lengths.
- `test_fwd_batched` confirms correctness when automatic `cu_seqlens` generation occurs for B > 1 fixed-length inputs.

## How to Run FlashKDA Correctness Tests

### Quick Sanity Run via Helper Script

For rapid verification during development, use the provided shell wrapper to install dependencies and execute the minimal test:

```bash
cd FlashKDA          # repository root

./tests/test.sh     # installs deps and executes tests/test_fwd.py

```

The script performs an editable install (`pip install -e .`), retrieves the required `flash-linear-attention` wheel, and launches `python tests/test_fwd.py`. Successful execution prints a confirmation message indicating exact match between the kernel and reference.

### Full Correctness Matrix Execution

For continuous integration or pre-release validation, execute the exhaustive suite:

```bash

# From the repository root

pip install -e .               # make the package importable

pip install "flash-linear-attention>=0.5.0" pytest   # ensure test deps

pytest tests/test_fwd_full.py -x -v          # sequential execution

```

The flags provide:
- **`-x`**: Abort immediately upon the first failure.
- **`-v`**: Verbose output displaying parametrised test IDs (e.g., `H96_T8192_state_dtype=bf16[in+out]`).

For faster execution across the thousand-plus test cases, enable parallel processing:

```bash
pytest tests/test_fwd_full.py -x -v -n auto

```

The `-n auto` option distributes test cases across available CPU cores, significantly reducing wall-clock time for the full matrix.

### Selective Test Execution

Target specific configurations using pytest’s keyword filtering against the generated test IDs:

```bash
pytest tests/test_fwd_full.py -k "H96 and T8192 and state_dtype=bf16 and in_only"

```

This command executes only the case where head count equals 96, sequence length equals 8192, bfloat16 state precision is used, and only the initial state is consumed.

## Summary

- **FlashKDA correctness tests** reside in [`tests/test_fwd.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/test_fwd.py) (minimal) and [`tests/test_fwd_full.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/test_fwd_full.py) (exhaustive).
- **Exact binary equality** is enforced via `torch.equal` against an FP64 `torch_ref` implementation.
- **Coverage includes** four state I/O modes, BF16/FP32 dtypes, head counts up to 96, sequences up to 1M tokens, and variable-length batching via `cu_seqlens`.
- **Execution methods** range from the [`./tests/test.sh`](https://github.com/MoonshotAI/FlashKDA/blob/main/./tests/test.sh) helper script for quick checks to `pytest -n auto` for parallelized full matrix validation.
- **Source files** implementing the logic include [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp) for the kernel and [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py) for the Python bindings.

## Frequently Asked Questions

### How does FlashKDA ensure numerical accuracy against the reference implementation?

FlashKDA guarantees accuracy by running the reference implementation in double precision (FP64) within [`torch_ref.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/torch_ref.py), then casting the results back to the target dtype (BF16 or FP32). The tests assert `torch.equal` between the CUDA kernel output and this reference, ensuring exact bitwise match rather than approximate tolerance, which is feasible because the mathematical operations are deterministic and the reference uses higher precision for the ground truth.

### What is the difference between [`test_fwd.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/test_fwd.py) and [`test_fwd_full.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/test_fwd_full.py)?

[`test_fwd.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/test_fwd.py) provides a minimal sanity check with fixed parameters for rapid developer feedback, verifying basic functionality on a single GPU. [`test_fwd_full.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/test_fwd_full.py) contains the exhaustive parametrized suite that systematically tests all combinations of state I/O modes, data types, head counts, sequence lengths (including extreme lengths), and batch configurations, making it suitable for CI pipelines and release validation.

### Can FlashKDA tests be executed without a GPU?

No. The correctness tests require a CUDA-capable GPU because they execute the native kernel compiled from [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp). The tests compare the GPU kernel output against a CPU-based PyTorch reference, so GPU hardware is mandatory for execution.

### How can I run only specific test configurations?

Use pytest’s `-k` flag with the test ID keywords generated by the parametrization. For example, `pytest tests/test_fwd_full.py -k "H32 and T1024"` runs only tests with 32 heads and 1024 sequence length across all other parameter combinations. This is useful for debugging specific failure modes without executing the entire matrix.