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

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 (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 handles the CHUNK=16 logic internally. Typical usage requires no manual chunk specification:

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:


# 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 (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, with implementation details in csrc/smxx/fwd_kernel1.cuh and csrc/smxx/fwd_launch.cu. The pure-Python reference in tests/torch_ref.py also validates these stability properties independently of the CUDA implementation.

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 →