> This page is for version Nightly (default).
> For other versions, use one of these documentation indexes:
> - Nightly (default): https://docs.nvidia.com/nemo/automodel/nightly/llms.txt
> - Latest: https://docs.nvidia.com/nemo/automodel/latest/llms.txt
> - 0.5.0 · 26.06: https://docs.nvidia.com/nemo/automodel/v0.5/llms.txt
> - 0.4.0 · 26.04: https://docs.nvidia.com/nemo/automodel/v0.4/llms.txt

> For clean Markdown of any page, append .md to the page URL.
> For a complete documentation index, see https://docs.nvidia.com/nemo/automodel/llms.txt.
> For AI client integration (Claude Code, Cursor, etc.), connect to the MCP server at https://docs.nvidia.com/nemo/automodel/_mcp/server.

# nemo_automodel.components.models.kimi_k3.situ_triton

Hand-written Triton kernels for the Kimi-K3 SiTU activation.

The SiTU chain (`beta * tanh(g / beta) * sigmoid(g) * linear_beta *
tanh(u / linear_beta) * w`) is a pure elementwise function of the gate/up
halves of one `[rows, 2 * intermediate]` projection plus one routing weight
per row. `torch.compile` (`BackendConfig.compile_situ`) fuses it into a
single kernel, but the generated 1-D pointwise kernels pay for two things on
every element: 64-bit `div`/`mod` to recover the row index for the
row-broadcast routing weight (the pointwise index space is flattened), and,
in the backward, both `cat` branches evaluated under `tl.where` (twelve
masked loads and four `tanh` per output element). On GB200 that leaves the
backward \~10x above its bandwidth bound (17 ms per 262k x 6144 bf16 call on a
2-node K3 profile).

These kernels tile the problem in 2-D (rows x columns) instead: the row index
is a per-tile 64-bit multiply, the column index stays 32-bit, each transcendental
is evaluated once, `d_gate` and `d_up` are stored straight into the two
halves of the output, and the routing-weight gradient (\`\`sum\_cols(go \* situ(g)

* up(u))`) is reduced inside the backward kernel while the tile is in
  registers. The fp32 math and its operation order are identical to
  `situ.\_situ\_fwd\_core`/`situ.\_situ\_bwd\_core\`\`; only the fp32 accumulation
  order of the routing-weight reduction differs.

## Module Contents

### Functions

| Name                                                                                         | Description                                                                                                   |
| -------------------------------------------------------------------------------------------- | ------------------------------------------------------------------------------------------------------------- |
| [`_check_2d_rows`](#nemo_automodel-components-models-kimi_k3-situ_triton-_check_2d_rows)     | -                                                                                                             |
| [`_rcp_fast`](#nemo_automodel-components-models-kimi_k3-situ_triton-_rcp_fast)               | `rcp.approx.ftz.f32` on an fp32 block of any shape (one SFU op; results below 2^-126 flush to zero).          |
| [`_situ_bwd_kernel`](#nemo_automodel-components-models-kimi_k3-situ_triton-_situ_bwd_kernel) | Analytic SiTU gradients for BLOCK\_R full rows; the routing-weight gradient is reduced in-kernel.             |
| [`_situ_fwd_kernel`](#nemo_automodel-components-models-kimi_k3-situ_triton-_situ_fwd_kernel) | out\[r, c] = beta \* tanh(g / beta) \* sigmoid(g) \* up(u) \* w\[r] for one (rows x cols) tile.               |
| [`_situ_sigmoid`](#nemo_automodel-components-models-kimi_k3-situ_triton-_situ_sigmoid)       | sigmoid of an fp32 block of any shape: SFU exp2 + approximate reciprocal when `FAST`, `tl.sigmoid` otherwise. |
| [`_situ_tanh`](#nemo_automodel-components-models-kimi_k3-situ_triton-_situ_tanh)             | tanh of an fp32 block of any shape: the SFU approximation when `FAST`, libdevice otherwise.                   |
| [`_tanh_fast`](#nemo_automodel-components-models-kimi_k3-situ_triton-_tanh_fast)             | MUFU.TANH (`tanh.approx.f32`, sm\_75+): one SFU op, \|rel err\| \<= 2^-11 — below bf16 resolution.            |
| [`_tile_configs`](#nemo_automodel-components-models-kimi_k3-situ_triton-_tile_configs)       | Autotune candidates: \~16 fp32 values per thread keeps the backward below the register cap.                   |
| [`situ_bwd_triton`](#nemo_automodel-components-models-kimi_k3-situ_triton-situ_bwd_triton)   | Weighted (or dense) SiTU backward.                                                                            |
| [`situ_fwd_triton`](#nemo_automodel-components-models-kimi_k3-situ_triton-situ_fwd_triton)   | Weighted (or dense) SiTU forward on `[rows, 2 * intermediate]` projections.                                   |

### Data

[`HAVE_TRITON`](#nemo_automodel-components-models-kimi_k3-situ_triton-HAVE_TRITON)

### API

```python
nemo_automodel.components.models.kimi_k3.situ_triton._check_2d_rows(
    name: str,
    t: torch.Tensor
) -> None
```

```python
nemo_automodel.components.models.kimi_k3.situ_triton._rcp_fast(
    x
)
```

`rcp.approx.ftz.f32` on an fp32 block of any shape (one SFU op; results below 2^-126 flush to zero).

```python
nemo_automodel.components.models.kimi_k3.situ_triton._situ_bwd_kernel(
    gu_ptr,
    rw_ptr,
    go_ptr,
    dgu_ptr,
    drw_ptr,
    n_rows,
    half,
    stride_gu,
    stride_go,
    stride_dgu,
    beta,
    linear_beta,
    HAS_RW: triton.language.constexpr,
    HAS_LINEAR: triton.language.constexpr,
    WANT_DRW: triton.language.constexpr,
    FAST_MATH: triton.language.constexpr,
    BLOCK_R: triton.language.constexpr,
    BLOCK_C: triton.language.constexpr
)
```

Analytic SiTU gradients for BLOCK\_R full rows; the routing-weight gradient is reduced in-kernel.

```python
nemo_automodel.components.models.kimi_k3.situ_triton._situ_fwd_kernel(
    gu_ptr,
    rw_ptr,
    out_ptr,
    n_rows,
    half,
    stride_gu,
    stride_out,
    beta,
    linear_beta,
    HAS_RW: triton.language.constexpr,
    HAS_LINEAR: triton.language.constexpr,
    FAST_MATH: triton.language.constexpr,
    BLOCK_R: triton.language.constexpr,
    BLOCK_C: triton.language.constexpr
)
```

out\[r, c] = beta \* tanh(g / beta) \* sigmoid(g) \* up(u) \* w\[r] for one (rows x cols) tile.

```python
nemo_automodel.components.models.kimi_k3.situ_triton._situ_sigmoid(
    x,
    FAST: triton.language.constexpr
)
```

sigmoid of an fp32 block of any shape: SFU exp2 + approximate reciprocal when `FAST`, `tl.sigmoid` otherwise.

```python
nemo_automodel.components.models.kimi_k3.situ_triton._situ_tanh(
    x,
    FAST: triton.language.constexpr
)
```

tanh of an fp32 block of any shape: the SFU approximation when `FAST`, libdevice otherwise.

```python
nemo_automodel.components.models.kimi_k3.situ_triton._tanh_fast(
    x
)
```

MUFU.TANH (`tanh.approx.f32`, sm\_75+): one SFU op, |rel err| \<= 2^-11 — below bf16 resolution.

```python
nemo_automodel.components.models.kimi_k3.situ_triton._tile_configs() -> list
```

Autotune candidates: \~16 fp32 values per thread keeps the backward below the register cap.

```python
nemo_automodel.components.models.kimi_k3.situ_triton.situ_bwd_triton(
    gate_up2: torch.Tensor,
    routing_weights2: torch.Tensor | None,
    grad_out2: torch.Tensor,
    beta: float,
    linear_beta: float | None,
    want_drw: bool,
    fast_math: bool = False
) -> tuple[torch.Tensor, torch.Tensor | None]
```

Weighted (or dense) SiTU backward.

**Parameters:**

**`gate_up2`** `torch.Tensor`

Saved gate+up projections of shape \[rows, 2 \* intermediate].

---

**`routing_weights2`** `torch.Tensor | None`

Saved routing weights of shape \[rows, 1], or None (dense).

---

**`grad_out2`** `torch.Tensor`

Upstream gradient of shape \[rows, intermediate], unit last stride.

---

**`beta`** `float`

SiTU beta applied to the gate branch.

---

**`linear_beta`** `float | None`

Optional bounded-linear beta applied to the up branch.

---

**`want_drw`** `bool`

Whether to reduce the routing-weight gradient (requires weights).

---

**`fast_math`** `bool` — default: False

Same SFU variants as in `situ_fwd_triton`; the autograd Functions in `situ.py` pass
the setting their forward used.

---

**Returns:** `torch.Tensor`

`(d_gate_up2, d_routing_weights2)` in the inputs' dtypes; the second

```python
nemo_automodel.components.models.kimi_k3.situ_triton.situ_fwd_triton(
    gate_up2: torch.Tensor,
    routing_weights2: torch.Tensor | None,
    beta: float,
    linear_beta: float | None,
    fast_math: bool = False
) -> torch.Tensor
```

Weighted (or dense) SiTU forward on `[rows, 2 * intermediate]` projections.

**Parameters:**

**`gate_up2`** `torch.Tensor`

Gate+up projections of shape \[rows, 2 \* intermediate] on a CUDA
device, unit stride along the last axis; gate in the first half.

---

**`routing_weights2`** `torch.Tensor | None`

Optional routing weights of shape \[rows, 1] (any float
dtype, contiguous), or None for the dense activation.

---

**`beta`** `float`

SiTU beta applied to the gate branch.

---

**`linear_beta`** `float | None`

Optional bounded-linear beta applied to the up branch.

---

**`fast_math`** `bool` — default: False

SFU `tanh.approx` / `exp2` / `rcp.approx` instead of libdevice tanh and an IEEE divide
(`KimiK3TextConfig.situ_backend = "triton_fast_math"`); at most one bf16 ulp from the exact chain.

---

**Returns:** `torch.Tensor`

Tensor of shape \[rows, intermediate] in `gate_up2`'s dtype.

```python
nemo_automodel.components.models.kimi_k3.situ_triton.HAVE_TRITON = True
```