FlashKDA Input Tensor Shapes and DTypes: Complete Specification for the Forward Pass
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 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], dtypetorch.bfloat16k(Key): Shape[B, T, H, K], dtypetorch.bfloat16v(Value): Shape[B, T, H, V], dtypetorch.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 (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], dtypetorch.bfloat16beta(Beta logits): Shape[B, T, H], dtypetorch.bfloat16. A sigmoid activation is applied internally to these logits.out(Output buffer): Shape[B, T, H, V], dtypetorch.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], dtypetorch.float32. Log-gate parameters applied per head.dt_bias: Shape[H, K], dtypetorch.float32. Bias terms for the gate computation.lower_bound: Pythonfloatconstraining 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 offlash_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). Acceptstorch.bfloat16ortorch.float32.final_state: Optional buffer receiving the last recurrent state. Must match the shape and dtype ofinitial_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], dtypetorch.int64. Cumulative sequence lengths tensor where entryicontains the starting index of sequencei. This mode requiresB = 1in 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:
- Device: All tensors must reside on CUDA (
device='cuda'). - Contiguity: Tensors must be contiguous in memory; non-contiguous views will raise an error.
- Dimensions: Head dimensions are fixed at
K = V = 128in 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:
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 and the validation tests in tests/test_fwd.py.
Summary
- Primary tensors (
q,k,v,g,beta,out) requiretorch.bfloat16, CUDA placement, and shapes following the[B, T, H, K|V]convention. - Gating parameters (
A_log,dt_bias) usetorch.float32for numerical precision. - State tensors support both
bfloat16andfloat32with shapes[B, H, V, K]or[N, H, V, K]depending on batching mode. - Critical constraint: Head dimensions are fixed at
K = V = 128as documented inflash_kda/__init__.py. - Variable-length processing requires
cu_seqlensof typeint64and batch sizeB = 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.
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 →