FlashKDA Delta Attention Algorithm: Linear-Time Recurrent Attention Explained

FlashKDA implements Kimi Delta Attention, a recurrent-style mechanism that replaces quadratic softmax attention with a linear-time Δ-rule recurrence operating on 16-token chunks.

The Delta Attention algorithm is the computational core of MoonshotAI/FlashKDA, designed to deliver the expressive power of transformer attention while maintaining O(T) complexity for long sequences. Unlike standard attention that computes pairwise token similarities across the entire sequence, FlashKDA processes input through a stateful recurrence that decays past information using learned gating parameters.

What is the Delta Attention Algorithm?

Delta Attention is a recurrent attention formulation that accumulates contextual information through a matrix state rather than materializing full attention maps. As implemented in csrc/flash_kda.cpp, the algorithm treats attention as a dynamic system where each token updates a hidden state matrix L of size D×D (where D = 128).

The mechanism relies on five input tensors: queries (q), keys (k), values (v), gate pre-activations (g), and decay factors (beta). These inputs undergo a two-kernel fusion strategy that splits computation between parallelizable per-chunk operations and sequential state updates.

Core Architecture of FlashKDA Delta Attention

Chunked Token Processing

The algorithm processes sequences in fixed chunks of 16 tokens (CHUNK = 16). Inputs arrive in shape [B, T, H, D] and are reshaped to a flat [T_total, H, D] layout in flash_kda/__init__.py, where T_total = B × T. This flattening enables Kernel K1 to compute quantities in parallel across all chunks, including:

  • Gate activation via sigmoid (implemented with tanh.approx.f32 PTX instruction)
  • L2-normalization of queries and keys
  • Decay factor computation A = exp(g_c * lower_bound) using base-2 exponentiation (ex2.approx.ftz.f32)
  • The chunk-wise matrix M_qk = q̃ᵀ·k̃ (a 16×16 matrix)

The Δ-Rule Recurrence State Update

Kernel K2 implements the core recurrence by walking chunk-by-chunk for each head. The algorithm updates the state matrix L according to the delta-rule:


# High-level pseudocode of the recurrence

for each chunk c:
    g_c = sigmoid(g_raw_c)          # Gate activation

    β_c = sigmoid(beta_c)           # Decay gating

    
    # Compute decayed projections

    q̃ = q_c * exp(g_c * lower_bound)
    k̃ = k_c * exp(g_c * lower_bound)
    
    # 16×16 outer product matrix

    M_qk = matmul(q̃, k̃)
    
    # Δ-rule state update with Neumann-series inversion

    L = (I - β_c · M_qk)⁻¹ · (L + g_c · (q̃ · v_cᵀ))

The matrix inverse (I - β_c·M_qk)⁻¹ is computed using a fp16 Neumann-series expansion, chosen because inverse entries remain within [-1, 1]. This recurrence continuously blends new value contributions into the state while decaying historical information via β_c.

Precision and Memory Optimizations

FlashKDA employs mixed precision to maximize throughput on NVIDIA GPUs:

  • State storage: The recurrent matrix L resides in bfloat16 on-chip, halving shared-memory usage compared to fp32.
  • Matrix inversion: Computed in fp16 as documented in docs/20260420-flashkda-v1-deep-dive.md.
  • Activations: Sigmoid gates leverage fast approximate PTX instructions (tanh.approx.f32).

Implementation Details in FlashKDA

Kernel K1: Per-Chunk Parallel Computation

The first CUDA kernel handles the embarrassingly parallel preprocessing. For every chunk, it computes decayed queries/keys and the M_qk matrices. This kernel operates on the reshaped tensor layout to maximize occupancy across streaming multiprocessors.

Kernel K2: Sequential State Recurrence

The second kernel enforces sequential dependencies between chunks. For each head, it maintains the D×D state matrix L in registers or shared memory, applying the Δ-rule update formula. The output for each token is derived by projecting the updated state onto the corresponding value vector.

State Management and Workspace Allocation

FlashKDA supports recurrent state passing for streaming inference scenarios. The optional initial_state and final_state tensors have shape [N, H, D, D], where N is the number of independent sequences.

The helper function get_workspace_size (defined in csrc/flash_kda.cpp) calculates the required temporary buffer size based on CHUNK, D, total tokens, heads, and sequence count. The function guarantees 128-byte alignment for all intermediate buffers to ensure coalesced memory access.

Code Example: Running FlashKDA Delta Attention

The following demonstrates a standard forward pass using the Python wrapper:

import torch
from flash_kda import fwd, get_workspace_size

# Configuration: Batch=1, Time=64, Heads=8, Dim=128

B, T, H, D = 1, 64, 8, 128

# Initialize bfloat16 inputs on CUDA

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)          # Gate pre-activation

beta = torch.randn(B, T, H, dtype=torch.bfloat16, device='cuda')

# Gate parameters

A_log = torch.randn(H, dtype=torch.float32, device='cuda')
dt_bias = torch.randn(H, D, dtype=torch.float32, device='cuda')
lower_bound = -3.0
scale = 1.0 / (D ** 0.5)

# Allocate output and workspace

out = torch.empty_like(q)
workspace_bytes = get_workspace_size(T, H)
workspace = torch.empty(workspace_bytes, dtype=torch.uint8, device='cuda')

# Execute Delta Attention forward pass

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

print(out.shape)   # torch.Size([1, 64, 8, 128])

For streaming inference with recurrent state:


# Initialize states for single sequence

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

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

# final_state now contains the updated recurrent state for next chunk

Summary

  • FlashKDA Delta Attention replaces quadratic softmax attention with a linear-time recurrence using a D×D state matrix updated via the Δ-rule.
  • The algorithm processes sequences in 16-token chunks with hidden dimension D=128, splitting work between a parallel preprocessing kernel (K1) and a sequential recurrence kernel (K2).
  • State updates follow L = (I - β·M_qk)⁻¹ · (L + g·(q·vᵀ)), with matrix inversion computed via fp16 Neumann-series expansion.
  • Mixed precision optimization stores states in bfloat16 while performing sensitive inversion operations in fp16.
  • Full recurrent state support enables streaming inference, with state tensors shaped [N, H, D, D] managed through optional parameters in flash_kda/__init__.py.

Frequently Asked Questions

How does Delta Attention differ from standard softmax attention?

Standard softmax attention computes pairwise token similarities across the entire sequence, resulting in O(T²) memory and computation. Delta Attention instead treats context accumulation as a recurrent dynamical system, updating a fixed-size state matrix L in O(T) time. This approach eliminates the need to materialize full attention maps while preserving the ability to model long-range dependencies through gated decay mechanisms.

What is the chunk size in FlashKDA and why 16 tokens?

FlashKDA processes input in chunks of 16 tokens (CHUNK = 16). This granularity balances parallel throughput with on-chip memory constraints. A chunk size of 16 allows efficient computation of the 16×16 M_qk matrices within shared memory while maintaining enough parallelism to saturate GPU streaming multiprocessors. The design document in docs/20260420-flashkda-v1-deep-dive.md confirms this choice optimizes the trade-off between memory bandwidth and compute utilization.

Why does FlashKDA use bfloat16 for state but fp16 for matrix inversion?

The recurrent state matrix L uses bfloat16 to reduce shared-memory footprint by 50% compared to fp32, enabling larger batch sizes and sequence lengths. However, the matrix inversion (I - β·M_qk)⁻¹ uses fp16 because the Neumann-series expansion requires values strictly within [-1, 1], where fp16 provides sufficient precision without the overhead of fp32. This mixed-precision strategy is documented in the FlashKDA deep-dive technical notes.

Can FlashKDA Delta Attention be used for streaming inference?

Yes. FlashKDA supports explicit recurrent state management through the initial_state and final_state parameters in the fwd function. By passing a state tensor of shape [N, H, D, D] and retrieving the updated state after the forward pass, applications can process arbitrarily long sequences in chunks without recomputing from the beginning. This capability is essential for real-time streaming applications where maintaining context across windows is required.

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 →