How FlashKDA Stores On-Chip Recurrent State in BF16 for Memory Efficiency

FlashKDA keeps the recurrent state directly in shared memory as bfloat16 (bf16), cutting the shared-memory footprint by roughly 50% compared to fp32 while performing updates using fp32 FMA instructions for numerical stability.

The FlashKDA kernel from the MoonshotAI/FlashKDA repository implements a memory-efficient recurrence mechanism for linear attention by storing the $V \times K$ state matrix in bf16 precision on-chip. This design eliminates costly fp32-to-bf16 conversions from the critical path of every recurrence step while preserving accuracy through mixed-precision computations. Understanding this storage strategy is essential for optimizing memory bandwidth and latency in high-performance linear attention implementations.

The BF16 State Layout in Shared Memory

FlashKDA organizes the recurrent state as a [N × H, D, D] tensor where $D$ represents the per-head dimension. Rather than allocating this buffer in fp32, the kernel defines a dedicated shared-memory layout called StateSmemLayout that holds bf16 elements.

According to the source code in csrc/smxx/fwd_kernel2.cuh (lines 82-99), this layout reserves space for the state_acc buffer—the primary accumulator where the recurrent state lives throughout kernel execution. By leveraging bf16's 16-bit representation, FlashKDA reduces the shared-memory requirement from 4 bytes per element to 2 bytes, effectively doubling the state capacity within the GPU's limited shared memory pool.

When the kernel initializes Tensor-Memory-Access (TMA) descriptors in csrc/smxx/fwd_launch.cu (lines 119-138), it distinguishes between bf16 and fp32 state formats. For bf16 states, the TMA loads data directly into the state_acc buffer. For fp32 inputs, the kernel allocates an intermediate state_fp32_buf to receive the data before conversion.

Loading States: Direct TMA vs. FP32 Fallback

The kernel handles initial state loading differently depending on the input precision, ensuring optimal memory access patterns while maintaining flexibility for different data types.

Direct BF16 Loading: When users supply an initial state already in bf16 format, FlashKDA performs a direct TMA load into the state_acc buffer. As implemented in csrc/smxx/fwd_kernel2.cuh (lines 242-259), this path bypasses any format conversion entirely, moving data straight from global memory to shared memory without intermediate registers or conversion overhead.

FP32 Conversion Path: If the initial state arrives in fp32 precision, the kernel first loads into the temporary state_fp32_buf, then converts each element to the BF16 type before writing to state_acc. This conversion logic appears in csrc/smxx/fwd_kernel2.cuh (lines 268-285) and utilizes conversion helpers defined in csrc/smxx/utils.cuh (such as bf16_to_f32 and make_fragment_like<BF16>). Crucially, this conversion occurs once per kernel launch during the loading phase, not during every recurrence iteration.

State Updates with FP32 FMA

Despite storing the state in bf16, FlashKDA maintains numerical accuracy by performing the actual recurrence computations in fp32. Inside the K2 kernel's recurrence loop, the state updates use fp32 FMA instructions (fp32_fma) to accumulate changes.

The computation flow works as follows:

  1. The kernel loads bf16 values from state_acc and converts them to fp32 for the arithmetic operation.
  2. The update calculation executes using full fp32 precision to prevent accumulation errors.
  3. The result rounds back to bf16 before writing to the shared-memory buffer.

This approach ensures that rounding to bf16 occurs only at the end of each update cycle, preserving numerical stability while maintaining the memory efficiency benefits of half-precision storage.

Storing the Final State

At kernel completion, FlashKDA writes the accumulated state back to global memory using optimized TMA store operations. The implementation in csrc/smxx/fwd_kernel2.cuh handles both output formats:

  • BF16 Output (lines 787-804): Executes a direct TMA store from the state_acc buffer to global memory without conversion overhead.
  • FP32 Output (lines 805-822): Performs an on-the-fly conversion from bf16 to fp32 before the TMA store, accommodating users who require full-precision final states.

This dual-path approach ensures that the internal bf16 representation never imposes precision constraints on the API boundary.

Implementation Example

The following Python example demonstrates how to leverage FlashKDA's bf16 on-chip storage when processing sequences:

import torch
import flash_kda

# Configure dimensions

B, T, H, D = 2, 128, 8, 64  # batch, sequence length, heads, per-head dim

K = D  # key dimension

# Create bf16 input tensors

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, K, 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')
scale = 1.0 / (D ** 0.5)

# Prepare output and state buffers

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

# Execute FlashKDA - state remains in bf16 on-chip throughout

flash_kda.fwd(
    q, k, v, g, beta, scale,
    out,
    initial_state=initial_state,
    final_state=final_state
)

# final_state contains the updated recurrent state in bf16

print(final_state.shape)  # (B, H, D, D)

As noted in the project's deep-dive documentation (docs/20260420-flashkda-v1-deep-dive.md, lines 47-50), this design "cuts the shared memory footprint of the state roughly in half and removes the fp32 → bf16 cast that would otherwise sit on the critical path of every bf16 GEMM feeding the state."

Summary

  • 50% Memory Reduction: Storing the recurrent state as bf16 in shared memory halves the storage requirement compared to fp32, enabling larger state matrices within hardware constraints.
  • Eliminated Conversion Overhead: By maintaining bf16 storage throughout the recurrence, FlashKDA removes per-iteration fp32-to-bf16 conversions from the critical execution path.
  • Mixed-Precision Accuracy: State updates use fp32 FMA operations internally, with rounding to bf16 occurring only after complete update cycles, preserving inference accuracy.
  • Flexible I/O: The kernel accepts both bf16 and fp32 initial states, converting once at load time, and outputs either format through dedicated TMA store paths.

Frequently Asked Questions

Why does FlashKDA use bf16 instead of fp16 for the recurrent state?

FlashKDA selects bf16 over fp16 because bf16 offers a larger dynamic range (same exponent bits as fp32) while still providing the 50% memory savings of 16-bit formats. This prevents overflow issues during state accumulation that could occur with fp16's limited range, particularly when processing long sequences where state values might grow large.

How does the bf16 storage affect numerical accuracy during inference?

The bf16 storage imposes minimal accuracy degradation because the actual recurrence arithmetic executes in fp32. State values convert to fp32 for FMA operations, accumulate with full precision, then round back to bf16 only when writing results. The FlashKDA benchmarks verify that this approach introduces no measurable accuracy loss compared to pure fp32 state storage.

Can I use fp32 states with FlashKDA if my application requires full precision?

Yes, FlashKDA supports fp32 state I/O through the initial_state and final_state parameters. When providing fp32 inputs, the kernel loads into a temporary buffer (state_fp32_buf), converts to bf16 for on-chip storage, computes using fp32 FMAs, and converts back to fp32 for the output. While the internal storage remains bf16 for efficiency, the API accommodates full-precision workflows.

Where is the state buffer physically located during kernel execution?

The recurrent state resides in shared memory (SRAM) on the GPU SM (Streaming Multiprocessor), allocated through the StateSmemLayout structure in csrc/smxx/fwd_kernel2.cuh. This placement eliminates global memory bandwidth bottlenecks during the recurrence computation, allowing the state to feed directly into the bf16 GEMM pipelines without intermediate global memory round-trips.

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 →