FlashKDA Correctness Tests: Validation Suite and Execution Guide

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

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:


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

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:

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 (minimal) and 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 helper script for quick checks to pytest -n auto for parallelized full matrix validation.
  • Source files implementing the logic include csrc/flash_kda.cpp for the kernel and 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, 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 and test_fwd_full.py?

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

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 →