# How FlashKDA Implements the Sigmoid Function Efficiently in its Gating Path

> Discover how FlashKDA efficiently implements the sigmoid function using a custom CUDA kernel and hardware acceleration for a single-cycle GPU operation, optimizing its gating path.

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

---

**FlashKDA uses a custom CUDA kernel that approximates the sigmoid function via the hardware-accelerated `tanh.approx.f32` instruction, avoiding costly exponential operations and reducing the gating path to a single-cycle GPU operation.**

FlashKDA replaces the standard PyTorch `torch.sigmoid` with a hardware-optimized approximation to accelerate the gating mechanism in its kernel fusion pipeline. This implementation leverages inline PTX assembly to compute `sigmoid(x) = 0.5 * tanh(0.5 * x) + 0.5` directly on the GPU, delivering significant latency reductions for the β (beta) activation and gate paths. According to the MoonshotAI/FlashKDA source code, this optimization is critical for maintaining throughput in the library's fused attention kernels.

## Custom CUDA Kernel with PTX Assembly

The efficient sigmoid implementation is defined in [`tests/torch_ref.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/torch_ref.py) as a raw CUDA kernel string that utilizes NVIDIA's hardware-accelerated hyperbolic tangent instruction.

The kernel employs the mathematical identity `sigmoid(x) = (tanh(x/2) + 1) / 2` to transform the exponential-based sigmoid into a hardware-friendly tanh operation:

```cpp
__global__ void sigmoid_tanh_fp32_kernel(const float* __restrict__ input,
                                         float* __restrict__ output, int n) {
    int idx = blockIdx.x * blockDim.x + threadIdx.x;
    if (idx < n) {
        float xh = input[idx] * 0.5f;
        float th;
        asm("tanh.approx.f32 %0, %1;" : "=f"(th) : "f"(xh));
        output[idx] = th * 0.5f + 0.5f;
    }
}

```

This inline assembly invokes `tanh.approx.f32`, a single-cycle approximation available on modern NVIDIA GPUs that bypasses the transcendental math unit's expensive exponential calculations.

## Runtime Compilation and Python Wrapper

Instead of pre-compiling the kernel, FlashKDA uses `torch.utils.cpp_extension.load_inline` to compile the CUDA source at import time. The Python wrapper exposes the kernel as `sigmoid_ext.sigmoid_tanh_fp32`:

```python
sigmoid_ext = load_inline(
    name='sigmoid_ext',
    cpp_sources='torch::Tensor sigmoid_tanh_fp32(torch::Tensor input);',
    cuda_sources=_sigmoid_cuda_src,
    functions=['sigmoid_tanh_fp32'],
)

```

The wrapper accepts `torch.float32` tensors and returns the activated output, maintaining full interoperability with PyTorch's autograd system while delivering native CUDA performance.

## Integration in the Gating Path

In the reference implementation, the custom sigmoid kernel replaces standard activations in two critical locations within [`tests/torch_ref.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/torch_ref.py).

### Gate Activation

At lines 64-66, the kernel scales the log-gate values before applying the sigmoid:

```python
g = scale * sigmoid_ext.sigmoid_tanh_fp32(a_log_exp * g)

```

### Beta Activation

At lines 14-16, per-head β logits are processed in float32 precision:

```python
beta_activated = sigmoid_ext.sigmoid_tanh_fp32(beta_chunk.to(torch.float32))

```

These calls ensure that both the gating mechanism and the β parameter use the hardware-accelerated path rather than PyTorch's default implementation.

## Why Hardware-Accelerated Tanh is Faster

The standard sigmoid implementation requires computing `1 / (1 + exp(-x))`, which involves an expensive exponential evaluation on the GPU's special function units. FlashKDA's approach offers three key advantages:

1. **Single-Cycle Execution:** The `tanh.approx.f32` instruction executes in one cycle on NVIDIA hardware, providing a low-latency approximation of the hyperbolic tangent.
2. **Avoided Exponential Costs:** By reformulating sigmoid in terms of tanh, the kernel eliminates the need for `expf()` calls, which typically require multiple cycles and higher energy consumption.
3. **Vectorized Throughput:** The element-wise kernel processes tensors with full CUDA thread parallelism, fitting naturally into FlashKDA's fused kernel pipeline without synchronization overhead.

## Using the Sigmoid in the Public API

When calling the high-level `flash_kda.fwd` function, users provide β as pre-activation logits. The library internally handles the sigmoid conversion via the optimized kernel, as documented in [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py):

```python
import torch
from flash_kda import fwd

B, T, H, D = 1, 32, 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)
beta = torch.randn(B, T, H, dtype=torch.bfloat16, device='cuda')
scale = 0.5
out = torch.empty_like(v)

fwd(q, k, v, g, beta, scale, out,
    A_log=torch.zeros(H, dtype=torch.float32, device='cuda'),
    dt_bias=torch.zeros(H, D, dtype=torch.float32, device='cuda'),
    lower_bound=-5.0)

```

For direct access outside the main kernel pipeline, instantiate the sigmoid extension from the reference tests:

```python
from tests.torch_ref import sigmoid_ext

x = torch.randn(1024, dtype=torch.float32, device='cuda')
result = sigmoid_ext.sigmoid_tanh_fp32(x)

```

## Summary

- FlashKDA implements sigmoid via `tanh.approx.f32` in [`tests/torch_ref.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/torch_ref.py) rather than using `torch.sigmoid`.
- The mathematical identity `sigmoid(x) = 0.5 * tanh(0.5 * x) + 0.5` enables single-cycle hardware approximation.
- The custom kernel is compiled at runtime using `torch.utils.cpp_extension.load_inline` and exposed as `sigmoid_ext.sigmoid_tanh_fp32`.
- Both gate activation (lines 64-66) and beta activation (lines 14-16) in the reference implementation use this optimized path.
- The public `flash_kda.fwd` API accepts pre-activation β logits and handles the efficient sigmoid conversion internally.

## Frequently Asked Questions

### Why does FlashKDA use tanh instead of the standard sigmoid formula?

FlashKDA uses the tanh-based reformulation because NVIDIA GPUs provide the `tanh.approx.f32` PTX instruction, which executes in a single cycle with hardware-level approximation. The standard sigmoid requires computing an exponential function, which consumes significantly more GPU cycles and energy. By calculating `sigmoid(x)` as `0.5 * tanh(0.5 * x) + 0.5`, the kernel achieves identical mathematical results with substantially lower latency.

### Where is the sigmoid kernel actually used in the FlashKDA codebase?

According to the source code in [`tests/torch_ref.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/tests/torch_ref.py), the `sigmoid_tanh_fp32` kernel is invoked at lines 14-16 for beta activation and lines 64-66 for gate activation. Additionally, the C++ entry point in [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp) forwards tensors to CUDA kernels that ultimately utilize this activation logic within the fused kernel pipeline.

### Can I use FlashKDA's efficient sigmoid in my own PyTorch code?

Yes. While the `flash_kda.fwd` function handles sigmoid activation internally for the β parameter, you can import the extension directly from `tests.torch_ref` to apply the hardware-accelerated sigmoid to arbitrary float32 tensors. The runtime-compiled `sigmoid_ext` module provides the `sigmoid_tanh_fp32` function for standalone use.

### What precision does the FlashKDA sigmoid kernel support?

The custom kernel operates specifically on `float32` (FP32) tensors, as indicated by the `tanh.approx.f32` instruction in the CUDA source. When processing β logits or gate values in other formats like bfloat16, the code explicitly casts inputs to float32 before calling `sigmoid_ext.sigmoid_tanh_fp32`, ensuring numerical stability while maintaining the performance benefits of the hardware approximation.