FE-OSS APIs Overview

View as Markdown

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:

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):

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:

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
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
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.

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.