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], 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 (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).

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:

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) 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.
  • 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.

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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →