# FlashKDA Dimension Constraints for K and V Tensors: Fixed 128‑Dimensional Requirement

> Understand FlashKDA dimension constraints for K and V tensors. Learn why the last dimension must be exactly 128 for optimal performance with MoonshotAI FlashKDA.

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

---

**FlashKDA strictly requires the key (`k`) and value (`v`) input tensors to have a fixed hidden dimension of exactly 128, meaning the last dimension of both tensors must equal `K = V = 128`.**

FlashKDA is a high‑performance CUDA implementation of additive attention mechanisms developed by MoonshotAI. Understanding the FlashKDA dimension constraints for K and V inputs is essential before integrating the kernel into your transformer architecture, as the underlying CUDA implementation hard‑codes a specific vector size that cannot be adjusted at runtime.

## Fixed Hidden Dimension Requirement

The `fwd` function in FlashKDA enforces a non‑negotiable constraint: the hidden dimension of keys and values must be exactly 128. This limitation is explicitly documented in the source code and strictly enforced by the underlying CUDA kernels.

According to the implementation in [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py), the expected tensor shapes are defined as:

- `k`: shape `[B, T, H, K]` (bf16) – where **K must be 128**【source:flash_kda/__init__.py#L9-L10】  
- `v`: shape `[B, T, H, V]` (bf16) – where **V must be 128**【source:flash_kda/__init__.py#L11-L12】

The docstring for the `fwd` function contains an explicit note: *"Currently requires `K = V = 128`"*【source:flash_kda/__init__.py#L29-L31】. This constraint exists because the custom CUDA kernels in `csrc/smxx/fwd_kernel1.cuh` and `csrc/smxx/fwd_kernel2.cuh` are hand‑optimized with hard‑coded 128‑dimensional vector operations. Any deviation from this size triggers a runtime assertion error in the kernel.

## Valid Input Shapes (K = V = 128)

The following example demonstrates the correct tensor initialization for FlashKDA inputs. All key and value tensors must have their last dimension set to 128.

```python
import torch
from flash_kda import fwd

B, T, H = 2, 64, 4          # batch size, sequence length, number of heads

K = V = 128                # required hidden dimension (fixed)

# Initialize tensors with bfloat16 precision on CUDA

q = torch.randn(B, T, H, K, dtype=torch.bfloat16, device="cuda")
k = torch.randn(B, T, H, K, dtype=torch.bfloat16, device="cuda")
v = torch.randn(B, T, H, V, dtype=torch.bfloat16, device="cuda")
g = torch.randn(B, T, H, K, dtype=torch.bfloat16, device="cuda")
beta = torch.randn(B, T, H, dtype=torch.bfloat16, device="cuda")

# Auxiliary parameters

scale = 1.0
out = torch.empty_like(v)
A_log = torch.randn(H, dtype=torch.float32, device="cuda")
dt_bias = torch.randn(H, K, dtype=torch.float32, device="cuda")
lower_bound = -2.0

# Optional recurrent states (None for standard mode)

initial_state = None
final_state = None
cu_seqlens = None

# Execute forward pass

fwd(q, k, v, g, beta, scale, out,
    A_log, dt_bias, lower_bound,
    initial_state, final_state, cu_seqlens)

```

## Invalid Dimensions and Runtime Errors

Passing tensors with hidden dimensions other than 128 results in immediate runtime failure. The CUDA kernel performs strict assertion checks on the vector size.

```python

# Incorrect configuration: K = V = 64

k_wrong = torch.randn(B, T, H, 64, dtype=torch.bfloat16, device="cuda")
v_wrong = torch.randn(B, T, H, 64, dtype=torch.bfloat16, device="cuda")

# This invocation will raise a runtime error

fwd(q, k_wrong, v_wrong, g, beta, scale, out,
    A_log, dt_bias, lower_bound)

```

The error originates from assertions within the CUDA kernel implementation, which validates that `K == V == 128` before executing the optimized attention computation.

## Implementation Details

The dimension constraint is enforced at multiple levels of the FlashKDA stack:

- **[`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py)**: Python wrapper that documents the `K = V = 128` requirement in function signatures and docstrings.
- **[`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp)**: C++ entry point that forwards arguments to the CUDA kernel without dimension transformation.
- **`csrc/smxx/fwd_kernel1.cuh` and `csrc/smxx/fwd_kernel2.cuh`**: CUDA kernel headers containing static assertions and loop unrolling specifically optimized for 128‑dimensional vectors.

Because the kernels utilize warp‑level primitives and shared memory layouts calculated for 128‑byte alignment (32 floats × 4 bytes), supporting arbitrary hidden dimensions would require significant kernel refactoring and recompilation.

## Summary

- **Fixed size mandate**: Both key and value tensors must have their last dimension exactly equal to 128 (`K = V = 128`).
- **Shape specification**: Expected tensor shapes are `[B, T, H, 128]` for both `k` and `v`, where B=batch, T=time/sequence, and H=heads.
- **Hard‑coded kernels**: The CUDA implementation in `csrc/smxx/` directories assumes 128‑dimensional vectors for memory access patterns and computation.
- **Runtime validation**: Passing non‑compliant dimensions triggers assertion errors in the CUDA runtime, not Python exceptions.

## Frequently Asked Questions

### What is the exact dimension requirement for K and V in FlashKDA?

FlashKDA requires both the key (`k`) and value (`v`) tensors to have a hidden dimension of exactly 128. Specifically, the shape parameters must be `K = 128` and `V = 128`, making the full tensor shapes `[B, T, H, 128]` for both inputs.

### Can I use FlashKDA with hidden dimensions other than 128?

No. The CUDA kernels in `csrc/smxx/fwd_kernel1.cuh` and `csrc/smxx/fwd_kernel2.cuh` are hard‑coded for 128‑dimensional vectors. To use different hidden sizes, you would need to modify the kernel source code, adjust the memory access patterns, and recompile the extension.

### What error occurs if I pass tensors with K ≠ 128 or V ≠ 128?

The FlashKDA `fwd` function will raise a runtime CUDA error originating from kernel assertions that validate `K == V == 128`. This occurs during the kernel launch in [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp), not during Python argument parsing.

### Are there plans to support variable hidden dimensions in FlashKDA?

The current implementation in [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py) explicitly states "Currently requires `K = V = 128`," suggesting future versions may relax this constraint. However, as of the latest source code, the 128‑dimensional requirement remains hard‑coded in the CUDA kernels with no runtime flexibility.