Transformer Engine documentation#

Warning

You are currently viewing unstable developer preview of the documentation. To see the documentation for the latest stable release, refer to:

Transformer Engine (TE) is an NVIDIA library for accelerating Transformer model training on NVIDIA GPUs. It combines optimized building blocks and fused kernels with automatic mixed-precision-style APIs for PyTorch and JAX, so low-precision training can be adopted without rewriting a training stack.

Transformer Engine manages the scaling factors, amax histories, and quantization metadata required by low-precision recipes. Its modules cover attention, linear layers, normalization, Mixture-of-Experts (MoE), and communication operations used in large-scale distributed training.

Highlights#

  • FP8 training on NVIDIA Hopper, Ada, Blackwell, and Rubin GPUs.

  • MXFP8 and NVFP4 training on NVIDIA Blackwell GPUs.

  • Optimized attention, GEMM, normalization, quantization, and fused Transformer and MoE modules.

  • PyTorch and JAX APIs with autocast-style contexts and configurable low-precision recipes.

  • Support for tensor, sequence, context, and EP, including communication overlap.

  • FP16 and BF16 optimizations on NVIDIA Ampere architecture GPUs and later.