FE-OSS APIs Overview
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 oncudnn.jax.call/ CuTeDSL’s nativecutlass.jaxbridge; seegemm_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_sm100entry point for each of those same families (built oncudnn.jax.call; each API page documents its exact jit contract). Contiguous grouped MXFP8 SwiGLU/dSwiGLU wrappers accept Torch tensors and canonical JAX arrays, including underjax.jit, with explicit FP8 output dtype andsf_vec_size=32for JAX. They dispatch JAX calls tocudnn.jax.grouped_gemm_swiglu/grouped_gemm_dswiglu; matchingcudnn.torchaliases 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 thejax.jit-compatiblegemm_proj_rope_mxfp8_jax_sm100entry point.
This folder documents the Python FE APIs implemented under python/cudnn. For details on currently implemented operations, see:
- Causal Conv1d and Decode Update
- FLA Integration Shims
- GEMM + Amax
- GEMM + RoPE + MXFP8 Projection
- Gated Attention Block (SM107) — 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
- Prepared BF16 Tail RoPE
- GEMM + SwiGLU
- GEMM + sReLU
- GEMM + dsReLU
- Grouped GEMM (BF16)
- Grouped GEMM + GLU (Unified)
- Grouped GEMM + GLU + Hadamard
- Grouped GEMM + GLU + Hadamard + Quant
- Grouped GEMM + dGLU (Unified)
- Grouped GEMM + SwiGLU (Legacy, Contiguous-only)
- Grouped GEMM + dSwiGLU (Legacy, Contiguous-only)
- Grouped GEMM + sReLU (Unified) — optionally tanh
soft-clamped via
tanh_clamp_scale - Grouped GEMM + dsReLU (Unified)
- Discrete Grouped GEMM + SwiGLU
- Discrete Grouped GEMM + dSwiGLU
- Grouped GEMM + Quant (Legacy, Dense-only)
- Grouped GEMM + Quant (Unified)
- Grouped GEMM + Wgrad
- Block Sparse Attention (BSA)
- DeepSeek Sparse Attention (DSA)
- Flex Attention and mask plan design
- HSTU Attention (Blackwell SM100/SM103)
- HSTU LayerNorm-Multiply-SiLU-Dropout (LMSD)
- Native Sparse Attention (NSA)
- CSA Fused Compressor
- DSv4.1 Vision RoPE Backward
- Engram Saved-State Gate
- RMSNorm + RHT + Amax
- SDPA Backward (SM120)
- NVFP4 Attention QAT Backward
- RMSNorm + SiLU
- DSv4.1 mHC projection/RMS backward
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[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:
(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
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
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_streamin class API,streamin 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, underwarn_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.
- Several kernels reduce with cross-CTA atomics, so those outputs are not bit-exact run to
run. Kernels react to
File structure and examples
- All FE OSS APIs are implemented in the
python/cudnndirectory. - Correctness tests/samples are implemented in the
test/python/fe_apidirectory.