How FlashKDA Manages Recurrent State Across Sequences and Batch Modes
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 exposes the fwd function, which handles tensor validation and workspace allocation before delegating to the raw kernel.
# 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:
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:
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:
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 thecu_seqlenscumulative length tensor. - The
flash_kda.fwdwrapper inflash_kda/__init__.pyhandles workspace allocation and delegates to_fwd_raw, supporting both processing modes through the optionalcu_seqlensparameter. - State layout transposition (enabled via
transpose_state_layoutflags 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. 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 and benchmarks/bench_fwd.py.
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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →