# How FlashKDA's CHUNK Size of 16 Improves Numerical Stability Over CHUNK 64

> Discover how FlashKDA's CHUNK size of 16 enhances numerical stability over CHUNK 64 by managing exponential terms and optimizing matrix inversion for improved performance.

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

---

**FlashKDA uses a CHUNK size of 16 instead of 64 to keep exponential terms within bf16 representable range, enable efficient 16×16 matrix inversion in fp16, and map cleanly to NVIDIA SM80 MMA instructions, eliminating the complex rescaling tricks required by larger chunk sizes.**

FlashKDA is a high-performance linear attention kernel developed by MoonshotAI that processes token sequences in fixed-size blocks. While previous Flash Linear-Attention implementations relied on `CHUNK = 64`, the FlashKDA codebase deliberately selects `CHUNK = 16` to maintain numerical stability during mixed-precision training. This design choice is hardcoded throughout the CUDA kernels and documented in the deep-dive technical notes.

## Why CHUNK Size Determines Numerical Stability

Processing attention recurrence in chunks requires computing exponential cumulative sums and matrix inversions within each block. The size of these chunks directly impacts whether intermediate values remain representable in lower-precision formats like bf16 and fp16 without expensive rescaling operations.

### bf16 Dynamic Range and Exponential Terms

According to [`docs/20260420-flashkda-v1-deep-dive.md`](https://github.com/MoonshotAI/FlashKDA/blob/main/docs/20260420-flashkda-v1-deep-dive.md) (lines 13-19), the gate values `g` are lower-bounded by `-5`. Inside each chunk, the kernel computes `exp(cumsum(g))`. With **CHUNK = 16**, the maximum possible exponent stays within the representable range of **bf16** (approximately ±65504), preventing overflow.

With **CHUNK = 64**, the cumulative sum over 64 tokens would push `exp(cumsum(g))` outside bf16 limits, requiring aggressive intra-chunk rescaling tricks to avoid numerical overflow or underflow.

### Efficient 16×16 Matrix Inversion

FlashKDA must invert a `CHUNK × CHUNK` matrix of the form `INV = I – L` during the forward pass. As noted in `csrc/smxx/fwd_kernel1.cuh` (lines 62-68), a **16×16 matrix inversion** can be computed efficiently using Neumann-series expansion and remains well-conditioned in **fp16**. The resulting inverse values lie in the range `[-1, 1]`, staying safely within the reduced dynamic range of fp16.

A **64×64 matrix inversion** would exhibit a larger condition number, higher computational cost, and increased risk of overflow, making it numerically fragile for mixed-precision kernels.

### SM80 MMA Hardware Mapping

The kernel launch configuration in `csrc/smxx/fwd_launch.cu` (lines 31-35) defines `constexpr int CHUNK = 16` and selects SM80-only MMA layouts. A 16-element tile maps cleanly onto NVIDIA's 8×8×4 MMA instructions, keeping the compute path short and avoiding fp32-to-bf16 cast operations that introduce rounding error.

A **CHUNK = 64** implementation would need to split tiles across multiple MMA operations, re-introducing extra casting and accumulation steps that degrade precision.

## Code Examples and Stability Verification

The public API in [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py) handles the CHUNK=16 logic internally. Typical usage requires no manual chunk specification:

```python
import torch
from flash_kda import fwd

# Configuration uses CHUNK = 16 internally

B, T, H, D = 1, 128, 8, 128
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 logits

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

# FlashKDA forward - CHUNK is fixed at 16 for stability

fwd(q, k, v, g, beta, scale=1.0, out=out,
    A_log=None, dt_bias=None, lower_bound=-5.0)

```

Attempting to emulate a larger `CHUNK = 64` design manually demonstrates the stability problem. The exponentials quickly overflow in bf16:

```python

# Dangerous: Emulating CHUNK = 64 leads to overflow

CHUNK = 64
g_chunk = g[..., :CHUNK]                     # 64-token slice

exp_cumsum = torch.exp(g_chunk.cumsum(dim=-2))  # May produce inf in bf16

print(exp_cumsum.isinf().any())   # → True for typical random g

```

The reference implementation in [`tests/torch_ref.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/torch_ref.py) (lines 48-55) mirrors the production kernel logic and explicitly sets `CHUNK = 16` to maintain these stability guarantees in pure PyTorch.

## Summary

- **FlashKDA's CHUNK = 16** keeps `exp(cumsum(g))` within bf16 representable range (±65504), avoiding the overflow that occurs with CHUNK = 64.
- **16×16 matrix inversion** stays numerically stable in fp16 with values bounded in `[-1, 1]`, while 64×64 inversion becomes ill-conditioned and expensive.
- The chunk size maps directly to **SM80 MMA instructions** without requiring extra precision-casting steps that degrade accuracy.
- These design choices eliminate the need for complex intra-chunk rescaling tricks required by larger chunk sizes in Flash Linear-Attention kernels.

## Frequently Asked Questions

### Why does CHUNK = 16 prevent bf16 overflow?

The gate values `g` have a lower bound of `-5`. Over a chunk of 16 tokens, the cumulative sum remains small enough that `exp(cumsum(g))` never exceeds bf16's maximum value of approximately 65504. With 64 tokens, the cumulative sum grows large enough to cause exponential overflow, requiring manual rescaling to prevent `inf` values.

### How does the 16×16 matrix inversion work?

FlashKDA computes the inverse of `I - L` using Neumann-series expansion. Because the 16×16 matrix size is small, the series converges quickly and all intermediate values remain in the `[-1, 1]` range, making fp16 computation safe and accurate without overflow risk.

### Can I change the CHUNK size in FlashKDA?

No. The value is hardcoded as `constexpr int CHUNK = 16` in `csrc/smxx/fwd_launch.cu` and throughout the CUDA kernels in `csrc/smxx/fwd_kernel1.cuh`. This constant is fundamental to the numerical stability guarantees and SM80 instruction mapping, and changing it would require modifying the source and rebuilding the extension.

### Where is the stability advantage documented?

The design rationale appears in [`docs/20260420-flashkda-v1-deep-dive.md`](https://github.com/MoonshotAI/FlashKDA/blob/main/docs/20260420-flashkda-v1-deep-dive.md), with implementation details in `csrc/smxx/fwd_kernel1.cuh` and `csrc/smxx/fwd_launch.cu`. The pure-Python reference in [`tests/torch_ref.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/torch_ref.py) also validates these stability properties independently of the CUDA implementation.