# How FlashKDA Manages Recurrent State Across Sequences and Batch Modes

> Discover how FlashKDA manages recurrent state across sequences and batch modes using a persistent tensor and cumulative length indicators. Learn from the MoonshotAI/FlashKDA repository.

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

---

**FlashKDA manages recurrent state by maintaining a persistent state tensor that the kernel updates in-place, supporting both batched mode for independent sequences and variable-length mode that resets state boundaries via cumulative sequence length indicators.** The implementation in the MoonshotAI/FlashKDA repository uses distinct state shapes and reset behaviors depending on whether sequences are processed as standard batches or packed variable-length tensors.

## State Management Architecture

FlashKDA implements **recurrent K-Delta Attention (KDA)** by carrying a state tensor forward from one token to the next. This tensor stores intermediate hidden-state-like values that the CUDA/CUTLASS kernel updates during the sequence scan. The management strategy bifurcates based on the presence of sequence boundary metadata.

### Batched Mode (cu_seqlens=None)

In standard **batched mode**, the state tensor expects shape `[B, H, V, K]`, where `B` is the batch size, `H` the number of heads, and `V`/`K` the value/key dimensions. Each batch element maintains an independent recurrent chain that never mixes with others. The kernel treats the first dimension as the batch axis and **never resets** states internally, allowing persistent memory across the full sequence length for every batch entry.

### Variable-Length Mode (cu_seqlens provided)

For **variable-length (packed) batches**, all sequences concatenate into a single tensor with `B==1`, and the state tensor reshapes to `[N, H, V, K]` for `N` logical sequences. The kernel receives a **cu_seqlens** tensor (int64, shape `[N+1]`) containing cumulative sequence lengths. When the scan reaches a sequence boundary indicated by `cu_seqlens`, the kernel **zero-initializes** the next sequence's state, effectively resetting the recurrent memory between independent sequences despite the contiguous memory layout.

## Python Wrapper Implementation

The public API in [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py) exposes the `fwd` function, which handles tensor validation and workspace allocation before delegating to the raw kernel.

```python

# flash_kda/__init__.py – core wrapper

def fwd(q, k, v, g, beta, scale, out, A_log, dt_bias,
        lower_bound, initial_state=None, final_state=None, cu_seqlens=None):
    B, T_seq, H = q.shape[0], q.shape[1], q.shape[2]
    T_total = B * T_seq
    N = cu_seqlens.numel() - 1 if cu_seqlens is not None else B

    workspace = torch.empty(get_workspace_size(T_total, H, N),
                            dtype=torch.uint8, device=q.device)

    _fwd_raw(q, k, v, g, beta, float(scale), out, workspace,
             A_log, dt_bias, lower_bound,
             initial_state=initial_state,
             final_state=final_state,
             cu_seqlens=cu_seqlens)

```

The wrapper calculates `T_total` as the total token count and `N` as the number of sequences, allocating temporary workspace accordingly. The underlying `_fwd_raw` kernel receives the state tensors and boundary information directly.

## Practical Code Examples

### Batched Forward with Persistent State

The following demonstrates standard batched processing where each batch element maintains its own recurrent state:

```python
import torch, flash_kda, math

B, T, H, D = 2, 1024, 96, 128
q = torch.randn(B, T, H, D, dtype=torch.bfloat16, device='cuda')
k = torch.randn_like(q)
v = torch.randn_like(q)
g = torch.randn_like(q)
beta = torch.randn(B, T, H, dtype=torch.bfloat16, device='cuda')
A_log = torch.rand(H, dtype=torch.float32, device='cuda')
dt_bias = torch.rand(H, D, dtype=torch.float32, device='cuda')
scale = 1.0 / math.sqrt(D)

# One recurrent state per batch element: [B, H, V, K]

init_state = torch.zeros(B, H, D, D, dtype=torch.bfloat16, device='cuda')
out = torch.empty_like(q)
final_state = torch.empty_like(init_state)

flash_kda.fwd(q, k, v, g, beta, scale, out,
              A_log=A_log, dt_bias=dt_bias,
              lower_bound=-5.0,
              initial_state=init_state,
              final_state=final_state)

```

Omitting `cu_seqlens` triggers batched mode, ensuring state persistence across the full `T` tokens for each of the `B` sequences.

### Variable-Length Packed Sequences

For mixed-length sequences packed into a single tensor, provide `cu_seqlens` to enable automatic state resetting:

```python
import torch, flash_kda, math

seq_lengths = [300, 700, 1024]
N = len(seq_lengths)
T_total = sum(seq_lengths)

# Packed tensor with B==1

q = torch.randn(1, T_total, 96, 128, dtype=torch.bfloat16, device='cuda')
k = torch.randn_like(q)
v = torch.randn_like(q)
g = torch.randn_like(q)
beta = torch.randn(1, T_total, 96, dtype=torch.bfloat16, device='cuda')
A_log = torch.rand(96, dtype=torch.float32, device='cuda')
dt_bias = torch.rand(96, 128, dtype=torch.float32, device='cuda')
scale = 1.0 / math.sqrt(128)

# Cumulative lengths: shape [N+1]

cu_seqlens = torch.tensor([0] + list(torch.cumsum(torch.tensor(seq_lengths), dim=0)),
                          dtype=torch.long, device='cuda')

# State per logical sequence: [N, H, V, K]

init_state = torch.zeros(N, 96, 128, 128, dtype=torch.bfloat16, device='cuda')
out = torch.empty_like(q)
final_state = torch.empty_like(init_state)

flash_kda.fwd(q, k, v, g, beta, scale, out,
              A_log=A_log, dt_bias=dt_bias,
              lower_bound=-5.0,
              initial_state=init_state,
              final_state=final_state,
              cu_seqlens=cu_seqlens)

```

The kernel reads `cu_seqlens` to identify boundaries between the 300, 700, and 1024 token sequences, resetting the recurrent state to zero at each boundary.

### Inspecting Final Recurrent State

After execution, `final_state` contains the hidden state at the last token of each sequence:

```python
print(final_state.shape)   # [B, H, V, K] for batched, [N, H, V, K] for packed

```

This tensor captures the recurrent memory at sequence termination, useful for continued generation or state inspection.

## Summary

- FlashKDA manages recurrent state through a **state tensor** updated in-place by the CUDA kernel during sequence scanning.
- **Batched mode** uses shape `[B, H, V, K]` with independent persistence across batch elements and no internal resets.
- **Variable-length mode** uses shape `[N, H, V, K]` with automatic state zeroing between sequences as indicated by the `cu_seqlens` cumulative length tensor.
- The `flash_kda.fwd` wrapper in [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py) handles workspace allocation and delegates to `_fwd_raw`, supporting both processing modes through the optional `cu_seqlens` parameter.
- State layout transposition (enabled via `transpose_state_layout` flags in tests and benchmarks) optimizes memory access for the underlying CUTLASS kernels.

## Frequently Asked Questions

### What is the shape of the recurrent state tensor in FlashKDA?

In batched mode, the state tensor has shape `[B, H, V, K]` where `B` is the batch size. In variable-length mode, the shape becomes `[N, H, V, K]` where `N` represents the number of packed sequences. The dimensions `H`, `V`, and `K` correspond to attention heads, value dimension, and key dimension respectively.

### How does FlashKDA reset the recurrent state between sequences?

The kernel resets state automatically only in variable-length mode when provided with a `cu_seqlens` tensor. Upon encountering the cumulative length index indicating a sequence boundary, the kernel zero-initializes the state for the subsequent sequence. In standard batched mode, states persist across the entire sequence length without reset.

### Can FlashKDA process mixed-length sequences in a single batch?

Yes, through **variable-length (packed) mode**. Concatenate sequences along the time dimension with `B==1`, provide the `cu_seqlens` tensor containing cumulative lengths, and the kernel will treat each logical sequence independently while resetting state at boundaries. This avoids padding overhead for heterogeneous sequence lengths.

### Where is the recurrent state updated in the FlashKDA codebase?

The recurrent state update occurs in the `_fwd_raw` kernel, called from [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py). The Python wrapper allocates workspace and passes `initial_state` and `final_state` tensors to this underlying CUDA/CUTLASS implementation. State layout transposition and boundary handling logic reside in the kernel implementation, with usage patterns demonstrated in [`tests/test_fwd.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/test_fwd.py) and [`benchmarks/bench_fwd.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/benchmarks/bench_fwd.py).