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

# FE-OSS APIs Overview 

**FE-OSS APIs are experimental and subject to change.**

The GEMM CuTeDSL APIs are type-erased and torch-lazy: torch is imported only when torch tensors are passed. JAX arrays are additionally accepted wherever the kernel's tensor layouts are expressible as row-major arrays (each API's page has a "JAX support" section with its exact contract):

- **Dense fusions** (amax, swiglu, srelu, dsrelu): full JAX eager support, plus `jax.jit`-compatible XLA custom-call entry points for all four (built on `cudnn.jax.call` / CuTeDSL's native `cutlass.jax` bridge; see `gemm_amax.md` "Using JAX arrays").
- **Grouped / discrete-grouped**: JAX eager support in discrete (pointer-array) weight modes — unfused grouped GEMM, glu/dglu (BF16), dsrelu (FP8), wgrad (BF16), and discrete-grouped swiglu/dswiglu (FP8) — plus a `jax.jit`-compatible `*_jax_sm100` entry point for each of those same families (built on `cudnn.jax.call`; each API page documents its exact jit contract). Contiguous grouped MXFP8 SwiGLU/dSwiGLU wrappers accept Torch tensors and canonical JAX arrays, including under `jax.jit`, with explicit FP8 output dtype and `sf_vec_size=32` for JAX. They dispatch JAX calls to `cudnn.jax.grouped_gemm_swiglu` / `grouped_gemm_dswiglu`; matching `cudnn.torch` aliases remain available. Existing Torch behavior and wrapper defaults are unchanged. Other dense weight modes, column-major bias layouts, and grouped srelu/quant, glu_hadamard, and block-scaled unified glu/dglu/wgrad backends reject JAX with clear errors.
- **proj_rope_mxfp8**: JAX eager support on both input paths with `w_out_in=True` (the transposed [in, out] weight view is torch-only), plus the `jax.jit`-compatible `gemm_proj_rope_mxfp8_jax_sm100` entry point.

This folder documents the Python FE APIs implemented under `python/cudnn`. For details on currently implemented operations, see:
- [Causal Conv1d](/cudnn/fe-oss-apis/causal_conv1d) and [Decode Update](/cudnn/fe-oss-apis/causal_conv1d_update)
- [FLA Integration Shims](/cudnn/fe-oss-apis/fla)
- [GEMM + Amax](/cudnn/fe-oss-apis/gemm_fusions/gemm_amax)
- [GEMM + RoPE + MXFP8 Projection](/cudnn/fe-oss-apis/gemm_fusions/gemm_proj_rope_mxfp8)
- [Gated Attention Block (SM107)](/cudnn/fe-oss-apis/gated_attention_block) — projection, QK-norm + RoPE, SDPA, sigmoid gate, out projection as one FROST block (bf16 / FP8 / MXFP8, optional MXFP4 weights and NVFP4 / MXFP4 output)
- [Tail RoPE + Microscaled QDQ](/cudnn/fe-oss-apis/rope-qdq)
- [Prepared BF16 Tail RoPE](/cudnn/fe-oss-apis/rope-tail)
- [GEMM + SwiGLU](/cudnn/fe-oss-apis/gemm_fusions/gemm_swiglu)
- [GEMM + sReLU](/cudnn/fe-oss-apis/gemm_fusions/gemm_srelu)
- [GEMM + dsReLU](/cudnn/fe-oss-apis/gemm_fusions/gemm_dsrelu)
- [Grouped GEMM (BF16)](/cudnn/fe-oss-apis/gemm_fusions/grouped_gemm)
- [Grouped GEMM + GLU (Unified)](/cudnn/fe-oss-apis/gemm_fusions/grouped_gemm_glu)
- [Grouped GEMM + GLU + Hadamard](/cudnn/fe-oss-apis/gemm_fusions/grouped_gemm_glu_hadamard)
- [Grouped GEMM + GLU + Hadamard + Quant](/cudnn/fe-oss-apis/gemm_fusions/grouped_gemm_glu_hadamard_quant)
- [Grouped GEMM + dGLU (Unified)](/cudnn/fe-oss-apis/gemm_fusions/grouped_gemm_dglu)
- [Grouped GEMM + SwiGLU (Legacy, Contiguous-only)](/cudnn/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu)
- [Grouped GEMM + dSwiGLU (Legacy, Contiguous-only)](/cudnn/fe-oss-apis/gemm_fusions/grouped_gemm_dswiglu)
- [Grouped GEMM + sReLU (Unified)](/cudnn/fe-oss-apis/gemm_fusions/grouped_gemm_srelu) — optionally tanh
  soft-clamped via `tanh_clamp_scale`
- [Grouped GEMM + dsReLU (Unified)](/cudnn/fe-oss-apis/gemm_fusions/grouped_gemm_dsrelu)
- [Discrete Grouped GEMM + SwiGLU](/cudnn/fe-oss-apis/gemm_fusions/discrete_grouped_gemm_swiglu)
- [Discrete Grouped GEMM + dSwiGLU](/cudnn/fe-oss-apis/gemm_fusions/discrete_grouped_gemm_dswiglu)
- [Grouped GEMM + Quant (Legacy, Dense-only)](/cudnn/fe-oss-apis/gemm_fusions/grouped_gemm_quant)
- [Grouped GEMM + Quant (Unified)](/cudnn/fe-oss-apis/gemm_fusions/grouped_gemm_quant_unified)
- [Grouped GEMM + Wgrad](/cudnn/fe-oss-apis/gemm_fusions/grouped_gemm_wgrad)
- [Block Sparse Attention (BSA)](/cudnn/fe-oss-apis/bsa)
- [DeepSeek Sparse Attention (DSA)](/cudnn/fe-oss-apis/dsa)
- [Flex Attention](/cudnn/fe-oss-apis/attention/flex_attention) and [mask plan design](/cudnn/fe-oss-apis/attention/flex_attention_design)
- [HSTU Attention (Blackwell SM100/SM103)](/cudnn/fe-oss-apis/hstu/hstu_attention)
- [HSTU LayerNorm-Multiply-SiLU-Dropout (LMSD)](/cudnn/fe-oss-apis/hstu/hstu_lmsd)
- [Native Sparse Attention (NSA)](/cudnn/fe-oss-apis/nsa)
- [CSA Fused Compressor](/cudnn/fe-oss-apis/csa)
- [DSv4.1 Vision RoPE Backward](attention/vision_rope_backward.md)
- [Engram Saved-State Gate](/fe-oss-apis/engram_saved_gate)
- [RMSNorm + RHT + Amax](/cudnn/fe-oss-apis/rmsnorm_rht_amax)
- [SDPA Backward (SM120)](/cudnn/fe-oss-apis/attention/sdpa_bwd_sm120)
- [NVFP4 Attention QAT Backward](/cudnn/fe-oss-apis/attention/nvfp4_attention_qat_backward)
- [RMSNorm + SiLU](/cudnn/fe-oss-apis/rmsnorm_silu)
- [DSv4.1 mHC projection/RMS backward](/cudnn/fe-oss-apis/gemm_fusions/mhc_projection_bwd)

## Installation and setup

All Frontend OSS APIs come installed with the `nvidia-cudnn-frontend` package, and so does the CuTeDSL runtime they JIT through — `nvidia-cutlass-dsl[cu13]>=4.6.2` and `apache-tvm-ffi` are required dependencies (and the DSL pulls `cuda-python` in transitively):
```bash
pip install nvidia-cudnn-frontend
```
(`pip install nvidia-cudnn-frontend[cutedsl]` still works; the `cutedsl` extra now names only `cuda-python`. A few APIs still want extras of their own — the cuTile linear-attention engines need `[cutile]`.)

The Triton NVFP4 attention QAT backward API additionally requires the
`triton` extra. Its API page documents the framework and GPU requirements.

Those required dependencies are framework-neutral. Install your tensor framework separately — from a checkout, the PEP 735 dependency groups pin the right companion packages:
```bash
pip install --group torch   # torch + torch-c-dlpack-ext
pip install --group jax     # jax >= 0.5 (XLA entry points via cutlass.jax, shipped with nvidia-cutlass-dsl)
```
(For the published wheel, `pip install torch torch-c-dlpack-ext` or `pip install "jax>=0.5"` directly.)

After installation, you can import the APIs directly from the `cudnn` package, i.e. `from cudnn import {your_operation}`

## API Usage

The causal-convolution family exposes semantic PyTorch operations only through
`cudnn.ops.causal_conv1d` and `cudnn.ops.causal_conv1d_update`; they return
ordinary tensors and keep prepared kernel objects private. The wrapper and
class conventions below apply to prepared frontend-only kernel APIs.

Most compiler-style operations expose the following two APIs. Functional
autograd integrations such as BSA and Flex Attention instead document their
own wrapper and reusable-plan lifecycle on their operation pages.

### 1. High-level wrapper

- Single pythonic function call
- Allocates and returns output tensors
- Returns outputs as a `TupleDict` (supports both dictionary-style key access and tuple unpacking)
- No explicit compilation step – internally caches compiled kernels via a simple dictionary lookup
- When to use:
  - Fast prototyping and common cases
  - You want automatic allocation and minimal boilerplate
  - You are okay with the library managing the compiled-kernel cache

```python
from cudnn import {your_operation}_wrapper
result = {your_operation}_wrapper(
    inputs,
    ...,
    config_options,
    ...,
    stream=None,
)

# Dictionary-style access (recommended)
primary_output = result["output_tensor_name"]

# Tuple unpacking (order follows documented wrapper output keys)
out0, out1 = result
```

### 2. Class API

- Explicit lifecycle with compile and execute steps
- Reusable object with underlying compiled kernel for multiple executions
- Requires preallocated output tensors
- When to use:
  - You need to reuse a compiled kernel across many calls
  - You want explicit control over compilation and lifecycle management
  
```python
from cudnn import {your_operation}

op = {your_operation}(
    sample_inputs,
    ...,
    sample_outputs,
    ...,
    config_options,
    ...
)
op.compile()
op.execute(
    inputs,
    ...
    outputs,
    ...
    current_stream=None,
)
```
Methods:
- `check_support()` – validates target problem configuration (i.e. tensor shapes, tensor strides, dtypes, tiling/cluster/kernel configurations, environment, etc.)
- `compile()` – compiles the kernel with the provided sample tensors and parameters.
- `execute(inputs, ..., outputs, ..., current_stream)` – runs the kernel with the provided inputs and outputs.
  
## Common Parameters and Conventions

- CUDA stream (`current_stream` in class API, `stream` in wrapper)
  - The cuda stream to use for operation kernel execution.
  - Default: None (uses default stream)

- Determinism
  - Several kernels reduce with cross-CTA atomics, so those outputs are not bit-exact run to
    run. Kernels react to `torch.use_deterministic_algorithms(True)`: they switch to a
    deterministic path where one exists, and otherwise raise (or warn, under
    `warn_only`) rather than silently return a non-reproducible result.
  - Where a deterministic path exists it can also be selected per call with
    `deterministic=True`, independent of the torch setting. See
    [Grouped GEMM + dsReLU](/cudnn/fe-oss-apis/gemm_fusions/grouped_gemm_dsrelu#deterministic-dprob).


## File structure and examples

- All FE OSS APIs are implemented in the `python/cudnn` directory.
- Correctness tests/samples are implemented in the `test/python/fe_api` directory.