# FlashKDA lower_bound Parameter: Controlling Exponentiation Precision in Selective State Space Models

> Learn how FlashKDA's lower_bound parameter controls exponentiation precision in selective state space models by setting minimum activation values for bf16 limits and fast exponential instructions.

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

---

**The `lower_bound` parameter in FlashKDA defines the minimum activation value for the input gate, which is converted to a scale factor (`gate_scale = lower_bound * LOG2E`) to bound exponentiation results within bf16 precision limits while enabling fast base-2 exponential instructions.**

FlashKDA is a high-performance CUDA implementation of linear attention mechanisms optimized for Mamba-style state space models. The `lower_bound` parameter plays a critical role in FlashKDA's exponentiation pipeline by regulating the numerical range of activation gates before they enter the cumulative sum operation, directly influencing the precision and stability of the forward pass.

## How lower_bound Shapes the Activation Gate

In FlashKDA, the input gate `g` undergoes a transformation pipeline that combines learned biases with bounded scaling. The `lower_bound` parameter serves as the floor for this transformation, ensuring that the activated gate cannot exceed negative thresholds that would destroy numerical precision.

### Conversion to the Gate Scale Factor

Before reaching the CUDA kernel, the Python-level `lower_bound` argument (typically set between **-5.0** and **0**) is transformed into a multiplicative scale factor in the C++ entry point. In [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp), the implementation multiplies the bound by the mathematical constant `LOG2E` (ln(2) inverse ≈ 1.4426950408889634):

```cpp
float gate_scale = float(lower_bound * 1.4426950408889634);   // Line 128

```

This `gate_scale` value is then passed to the device kernel where it modulates the post-activation gate values. Inside `csrc/smxx/fwd_kernel1.cuh`, the kernel applies this scale after the sigmoid/tanh approximation:

```cpp
g_val = gate_scale * sigmoid_tanh_approx_f32(g_val);        // Line 22

```

This multiplication remaps the natural logarithm domain of the gate to base-2 logarithms, preparing the values for the `ex2.approx.ftz.f32` hardware instruction.

## Numerical Precision Safeguards

The `lower_bound` parameter enforces two complementary precision guarantees that prevent catastrophic overflow in low-precision tensor formats.

### Bounding the bf16 Representable Range

By limiting how negative the pre-activation gate can become, `lower_bound` ensures that the cumulative sum `cumsum(g)` remains within the narrow dynamic range of **bfloat16** (bf16). According to the FlashKDA deep-dive documentation, with a configuration of `lower_bound = -5.0` and `CHUNK = 16`, the values of `exp(cumsum(g))` stay safely inside bf16’s representable range. Without this constraint, the exponential of large negative accumulations would underflow to zero or require expensive rescaling operations that reduce throughput.

### Optimizing for Base-2 Exponentiation

FlashKDA kernels utilize the `ex2.approx.ftz.f32` PTX instruction for maximum throughput, which computes base-2 exponentials rather than natural exponentials. The pre-multiplication by `LOG2E` in the `gate_scale` calculation effectively converts the natural-log-based `lower_bound` into a base-2 logarithmic scale factor. This removes the need for an explicit change-of-base division inside the performance-critical kernel loop, allowing the hardware to execute a single fast exponentiation without additional arithmetic overhead.

## Implementation in the FlashKDA Source Code

The precision control flows through three architectural layers:

1. **Python Interface** ([`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py)): Accepts the `lower_bound` float and forwards it to the C++ binding.
2. **C++ Bridge** ([`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp)): Transforms the bound into `gate_scale` using `LOG2E` as shown above.
3. **CUDA Kernel** (`csrc/smxx/fwd_kernel1.cuh`): Applies the scale to the activated gate values before the cumulative scan and exponentiation steps.

This pipeline ensures that the constraint defined at the Python API level propagates directly to the PTX instruction selection, maintaining bit-exact precision semantics across the software stack.

## Practical Configuration Example

When invoking the forward pass, set `lower_bound` to a negative value that balances expressiveness with numerical safety. A value of **-5.0** is empirically validated to preserve bf16 precision across sequence lengths up to the tested chunk sizes:

```python
import torch
import flash_kda

B, T, H, D = 2, 128, 12, 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)               # pre-activation gate

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

# Learned parameters

A_log = torch.randn(H, dtype=torch.float32, device='cuda')
dt_bias = torch.randn(H, D, dtype=torch.float32, device='cuda')

# Precision control: -5.0 keeps exp(cumsum(g)) within bf16 range

lower_bound = -5.0
out = torch.empty_like(v)

flash_kda.fwd(q, k, v, g, beta, 1.0, out, A_log, dt_bias, lower_bound)

```

Reducing the magnitude of `lower_bound` (e.g., to **-1.0**) allows larger gate activations that can push `exp(cumsum(g))` beyond the upper limits of bf16 representation, leading to infinity values or gradient instability during training.

## Summary

- The `lower_bound` parameter defines the minimum pre-activation value for FlashKDA's input gate, typically set between **-5.0** and **0**.
- It is converted to `gate_scale` via multiplication by `LOG2E` in [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp) to align with base-2 exponentiation instructions.
- This scaling bounds the cumulative sum of gates to remain within the **bf16** representable range, preventing underflow/overflow without runtime rescaling.
- The implementation spans the Python API, C++ bridge, and CUDA kernel layers to ensure consistent precision control.

## Frequently Asked Questions

### What happens if lower_bound is set too close to zero?

Setting `lower_bound` near zero (e.g., **-0.1**) restricts the gate to small negative values, which limits the dynamic range of the state transitions and may reduce model capacity. More critically, if the cumulative sum of gates grows too large in the positive direction, `exp(cumsum(g))` can exceed the **bf16** maximum representable value (~3.4 × 10³⁸), resulting in numerical overflow and NaN gradients during backpropagation.

### Why does FlashKDA use LOG2E to scale lower_bound?

FlashKDA uses the CUDA `ex2.approx.ftz.f32` instruction for fast base-2 exponentiation. Multiplying `lower_bound` by **LOG2E** (≈1.4427) converts the natural logarithm domain of the gate into base-2 logarithms. This eliminates the need for an expensive division or multiplication inside the kernel to change the exponent base, directly feeding the hardware-optimized instruction with correctly scaled inputs.

### How does lower_bound interact with dt_bias?

The `dt_bias` parameter provides a learnable shift to the raw gate values before the sigmoid/tanh activation, while `lower_bound` sets the hard floor after activation. The operation flow is: `g_activated = sigmoid(g_raw + dt_bias)`, then `g_scaled = lower_bound * LOG2E * g_activated`. Thus, `dt_bias` controls the centering of the gate distribution, and `lower_bound` controls the minimum magnitude of the scaled output that enters the exponential cumulative sum.

### Can lower_bound be positive?

While the API accepts positive values, setting `lower_bound > 0` would reverse the intended numerical protection. A positive bound would multiply the activated gate (range 0 to 1) by a positive `gate_scale`, producing only non-negative cumulative sums that grow exponentially without the stabilizing decay mechanism that negative bounds provide. This would quickly overflow bf16 ranges and destabilize training.