# FlashKDA Input Tensor Shapes and DTypes: Complete Specification for the Forward Pass

> Understand FlashKDA input tensor shapes and dtypes for the forward pass. Learn requirements for attention and parameter tensors on CUDA.

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

---

**FlashKDA's `fwd` function requires `bfloat16` tensors for attention inputs with shapes `[B, T, H, K]` or `[B, T, H, V]`, alongside `float32` parameter tensors for gating variables, with all tensors residing on CUDA devices and adhering to the fixed head dimension constraint `K = V = 128`.**

FlashKDA is a high-performance CUDA implementation of efficient linear attention mechanisms maintained by MoonshotAI. The public Python API exposed in [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py) centers on the `fwd` function, which serves as a thin validation wrapper around the compiled CUDA kernel before dispatch. Correct usage depends on satisfying strict contracts for tensor shapes, data types, and memory layout that are enforced at the Python boundary.

## Core Input Tensors for FlashKDA

The primary attention computation consumes five input tensors and one output buffer, all restricted to `torch.bfloat16` and CUDA memory.

### Query, Key, and Value Tensors

The attention mechanism operates on standard multi-head projections with batch-major layout:

- **`q`** (Query): Shape `[B, T, H, K]`, dtype `torch.bfloat16`
- **`k`** (Key): Shape `[B, T, H, K]`, dtype `torch.bfloat16`  
- **`v`** (Value): Shape `[B, T, H, V]`, dtype `torch.bfloat16`

Where `B` is batch size, `T` is sequence length, `H` is the number of heads, and `K`/`V` are head dimensions. According to the source docstring in [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py) (lines 9-12), the kernel currently enforces `K = V = 128`.

### Gating Tensors and Output Buffer

FlashKDA utilizes input-dependent gating mechanisms requiring additional `bfloat16` tensors:

- **`g`** (Gate pre-activation): Shape `[B, T, H, K]`, dtype `torch.bfloat16`
- **`beta`** (Beta logits): Shape `[B, T, H]`, dtype `torch.bfloat16`. A sigmoid activation is applied internally to these logits.
- **`out`** (Output buffer): Shape `[B, T, H, V]`, dtype `torch.bfloat16`. This tensor is modified in-place by the kernel.

### Global Scaling Factor

The `scale` parameter is passed as a **Python `float`**, not a tensor. It applies a global multiplicative factor to the attention scores before the gating mechanism processes them.

## Gating Parameters and Constraints

Separate from the attention inputs, FlashKDA requires learned or configured parameters that control the data-dependent gating dynamics. These use higher-precision `float32` for numerical stability:

- **`A_log`**: Shape `[H]`, dtype `torch.float32`. Log-gate parameters applied per head.
- **`dt_bias`**: Shape `[H, K]`, dtype `torch.float32`. Bias terms for the gate computation.
- **`lower_bound`**: Python `float` constraining the gate lower bound. The implementation requires this value to reside in the range `[-5.0, 0.0]` as validated in the wrapper logic (lines 18-20 of [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py)).

## Optional Recurrent States and Variable-Length Sequences

FlashKDA supports recurrent processing modes through optional state tensors and variable-length batching:

### Initial and Final States

For recurrent or chunked computation, you may supply state buffers:

- **`initial_state`**: Optional tensor with shape `[B, H, V, K]` for batched mode or `[N, H, V, K]` for variable-length mode (`N` = number of sequences). Accepts `torch.bfloat16` or `torch.float32`.
- **`final_state`**: Optional buffer receiving the last recurrent state. Must match the shape and dtype of `initial_state`.

These states represent the recurrent hidden state carried across sequence chunks, implemented according to the contract documented in lines 21-23 of the wrapper.

### Variable-Length Mode

When processing concatenated variable-length sequences (e.g., for packed training data):

- **`cu_seqlens`**: Shape `[N+1]`, dtype `torch.int64`. Cumulative sequence lengths tensor where entry `i` contains the starting index of sequence `i`. This mode requires `B = 1` in the input tensors, as the batch dimension is implicitly defined by the cumulative lengths.

## Memory Layout and Hardware Requirements

All input tensors must satisfy three hard constraints enforced before CUDA dispatch:

1. **Device**: All tensors must reside on CUDA (`device='cuda'`).
2. **Contiguity**: Tensors must be contiguous in memory; non-contiguous views will raise an error.
3. **Dimensions**: Head dimensions are fixed at `K = V = 128` in the current implementation.

## Complete Forward Pass Example

The following runnable example constructs valid inputs for a batch of 2 sequences with 128 time steps, 8 heads, and 128-dimensional keys/values:

```python
import torch
from flash_kda import fwd

B, T, H, K, V = 2, 128, 8, 128, 128  # K and V must equal 128

# Required bfloat16 attention inputs 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')
out = torch.empty_like(v)  # [B, T, H, V]

# Float32 gating parameters

A_log = torch.randn(H, dtype=torch.float32, device='cuda')
dt_bias = torch.randn(H, K, dtype=torch.float32, device='cuda')

# Python float constants

scale = 1.0
lower_bound = -2.0  # Must be within [-5.0, 0.0]

# Optional recurrent states

initial_state = torch.zeros(B, H, V, K, dtype=torch.bfloat16, device='cuda')
final_state = torch.empty_like(initial_state)

# Execute kernel

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

assert out.shape == torch.Size([B, T, H, V])

```

This example mirrors the construction logic found in [`benchmarks/bench_fwd.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/benchmarks/bench_fwd.py) and the validation tests in [`tests/test_fwd.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/test_fwd.py).

## Summary

- **Primary tensors** (`q`, `k`, `v`, `g`, `beta`, `out`) require `torch.bfloat16`, CUDA placement, and shapes following the `[B, T, H, K|V]` convention.
- **Gating parameters** (`A_log`, `dt_bias`) use `torch.float32` for numerical precision.
- **State tensors** support both `bfloat16` and `float32` with shapes `[B, H, V, K]` or `[N, H, V, K]` depending on batching mode.
- **Critical constraint**: Head dimensions are fixed at `K = V = 128` as documented in [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py).
- **Variable-length** processing requires `cu_seqlens` of type `int64` and batch size `B = 1`.

## Frequently Asked Questions

### What data types does FlashKDA support for input tensors?

FlashKDA enforces `torch.bfloat16` for all attention-related tensors including queries, keys, values, gates, and the output buffer. Gating parameters (`A_log`, `dt_bias`) must be `torch.float32`. Optional recurrent states (`initial_state`, `final_state`) may use either `torch.bfloat16` or `torch.float32`. The `cu_seqlens` tensor for variable-length mode must be `torch.int64`.

### What are the exact shape requirements for FlashKDA's recurrent state?

Recurrent states must have shape `[B, H, V, K]` when using standard batching, or `[N, H, V, K]` when processing `N` variable-length sequences with the `cu_seqlens` parameter. The head dimensions are constrained to `K = V = 128` in the current kernel implementation.

### Does FlashKDA support variable-length sequence batches?

Yes. FlashKDA supports variable-length sequences through the optional `cu_seqlens` argument, which accepts a `torch.int64` tensor of shape `[N+1]` containing cumulative sequence lengths. When using this mode, the batch dimension `B` in the primary input tensors must be set to 1, as the batch is implicitly defined by the cumulative lengths array.

### Are there hardware or memory layout constraints for FlashKDA inputs?

All tensors must reside on CUDA devices and be contiguous in memory. The kernel does not support strided or non-contiguous tensor layouts. Additionally, the `lower_bound` parameter must be a Python float within the range `[-5.0, 0.0]` to satisfy the gate activation constraints enforced in [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py).