# FlashKDA State Update Numerical Precision: How FP32 FMA Instructions Work in the Reference Implementation

> Discover how FlashKDA uses FP32 FMA instructions for state updates, internally casting to float64 for precision and back to float32 for efficiency. Learn about deterministic rounding and computational accuracy in the reference ...

- Repository: [Moonshot AI/FlashKDA](https://github.com/MoonshotAI/FlashKDA)
- Tags: internals
- Published: 2026-07-31

---

**FlashKDA performs recurrent state updates using FP32 FMA (fused multiply-add) instructions that internally cast operands to float64 for the arithmetic operation, then cast the result back to float32 before storage, ensuring deterministic rounding while maintaining computational accuracy.**

FlashKDA, developed by MoonshotAI, implements efficient linear attention mechanisms that demand careful numerical precision management during state updates. Understanding the exact numerical precision used for FlashKDA's state update with fp32 FMA instructions is essential for reproducing results and debugging numerical stability. The reference implementation in [`tests/torch_ref.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/torch_ref.py) reveals a specific precision pattern that balances performance and accuracy through strategic type promotion.

## Understanding FP32 FMA State Updates in FlashKDA

The state update mechanism relies on a custom helper function that simulates hardware-level fused multiply-add behavior while avoiding precision loss. Rather than performing the multiply-add directly in float32, the implementation temporarily promotes tensors to float64 for the intermediate computation.

This approach ensures that the multiplication and addition operations occur with higher precision before the final rounding to float32. The result is a deterministic FP32 FMA operation that matches hardware behavior on modern GPUs and TPUs, preventing cumulative rounding errors during long sequences.

### The fp32_fma Helper Function Implementation

In [`tests/torch_ref.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/torch_ref.py) (lines 94-99), the `fp32_fma` function implements the numerical precision pattern:

```python
def fp32_fma(c, a, b):
    return (c.to(torch.float64) + a.to(torch.float64) * b.to(torch.float64)).to(torch.float32)

```

This function accepts three float32 tensors—`c` (the accumulator), `a`, and `b` (the multiplicands)—casts them to float64 for the fused multiply-add operation, then explicitly casts the result back to float32. This sequence guarantees that rounding occurs exactly once, after the full-precision intermediate calculation completes.

### Integration with the Recurrent State Update

The actual state update occurs during chunk processing in the reference implementation. At lines 242-244 of [`tests/torch_ref.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/torch_ref.py), the code applies this FP32 FMA operation to update the recurrent state:

```python
work_state[seq_idx, h] = fp32_fma(delta_s,
                                   state_slice.to(torch.float32).t(),
                                   g_total_exp).to(torch.bfloat16).t()

```

Here, `delta_s` serves as the accumulator, while `state_slice` and `g_total_exp` provide the multiplicands. The result undergoes an additional cast to bfloat16 for storage in the work state tensor, demonstrating how FlashKDA maintains FP32 numerical precision for the computation while potentially using lower precision for memory storage.

## Practical Example of FP32 FMA State Updates

To implement the same numerical precision pattern in your own FlashKDA integration, use the following approach:

```python
import torch
from tests.torch_ref import fp32_fma

# Initialize FP32 tensors representing state and inputs

c = torch.randn(4, 4, dtype=torch.float32)  # Previous state

a = torch.randn(4, 4, dtype=torch.float32)  # Input weights

b = torch.randn(4, 4, dtype=torch.float32)  # Gate values

# Perform FP32 state update with internal FP64 precision

new_state = fp32_fma(c, a, b)
assert new_state.dtype == torch.float32

```

This implementation matches the numerical behavior found in the FlashKDA reference tests, ensuring that your custom kernels produce identical results to the PyTorch reference implementation when using FP32 FMA instructions for state updates.

## Summary

- **FlashKDA uses FP32 numerical precision** for state updates, implemented through a custom FMA helper that matches hardware behavior.
- **Intermediate computations use float64** to prevent precision loss during the multiply-add operation, with explicit casting back to float32 for the final result.
- **The `fp32_fma` function** in [`tests/torch_ref.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/torch_ref.py) (lines 94-99) centralizes this logic, making the numerical precision strategy explicit and testable.
- **State tensors may be stored in bfloat16**, but the update computation itself maintains full FP32 precision through the FMA operation.

## Frequently Asked Questions

### Why does FlashKDA use float64 for intermediate FMA calculations?

The temporary promotion to float64 prevents catastrophic cancellation and rounding errors that can accumulate during recurrent state updates over long sequences. By performing the multiply-add in higher precision and rounding only once at the end, FlashKDA ensures numerical stability while keeping the final state representation compact in float32 or bfloat16.

### Is the final recurrent state stored in FP32 or a different precision?

While the FP32 FMA operation produces a float32 result, the reference implementation at lines 242-244 of [`tests/torch_ref.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/torch_ref.py) shows the result being cast to `torch.bfloat16` before storage in `work_state`. This means the computation uses FP32 precision, but the memory representation uses the more compact bfloat16 format for efficiency.

### How does this numerical precision affect FlashKDA's reproducibility?

The explicit `fp32_fma` helper ensures deterministic behavior across different hardware platforms. Because it controls the exact point of rounding (after the FP64 computation), the state update produces identical results whether running on CPUs, GPUs, or TPUs, eliminating precision-related divergence between reference and kernel implementations.

### Can I modify the numerical precision for custom FlashKDA implementations?

While you can adjust the dtype parameters in `fp32_fma`, modifying the intermediate precision from float64 to float32 would change the numerical behavior and potentially break compatibility with the reference tests. The current implementation reflects a careful balance between accuracy and performance confirmed by the MoonshotAI test suite.