About this project
Transformer Engine (TE) is an NVIDIA library for accelerating Transformer models on NVIDIA GPUs. Its central feature is low-precision computation: 8-bit floating point (FP8) on Hopper, Ada and Blackwell GPUs, plus MXFP8 and NVFP4 formats on Blackwell, intended to improve performance and reduce memory use in both training and inference. It also supports optimizations across FP16 and BF16 on Ampere and later architectures.
The library provides highly optimized building blocks for common Transformer architectures and an automatic mixed-precision-like API that can be used with framework-specific code. A framework-agnostic C++ API is included so other deep learning libraries can add FP8 support for Transformers. TE modules internally maintain scaling factors and related values needed for FP8 training, which simplifies mixed-precision workflows for users.
Highlights listed by the project include easy-to-use modules for building Transformer layers with FP8 support, fused-kernel optimizations, FP8 support on Hopper, Ada and Blackwell, MXFP8 and NVFP4 support on Blackwell, and optimizations across FP16/BF16 on Ampere and later.
Usage examples are given for PyTorch and JAX/Flax. In PyTorch, users import transformer_engine.pytorch, create a recipe such as DelayedScaling with a chosen FP8 format, and wrap the forward pass in te.autocast. In JAX/Flax, similar autocast usage is shown with te_flax modules and a recipe. A getting-started guide is linked for a fuller tutorial.
Installation options include NGC Docker containers (recommended), pip packages with extras for PyTorch and/or JAX, conda-forge packages for the PyTorch integration, and source builds. System requirements mention Blackwell, Hopper, Grace Hopper/Blackwell, Ada and Ampere hardware; Linux as the official OS with limited WSL2 support; CUDA 12.1+ (12.8+ for Blackwell); cuDNN 9.12+; GCC 9+ or Clang 10+ with C++17; and Python 3.12 recommended. FP8 features require compute capability 8.9 or higher. Build-related environment variables such as CUDA_PATH, CUDNN_PATH, CXX, NVTE_FRAMEWORK, MAX_JOBS and NVTE_CUDA_ARCHS are documented.
The README covers FlashAttention-2 and FlashAttention-3 support in PyTorch, with FlashAttention-3 prioritized when both are present, and notes that FlashAttention-2 compilation can be memory-intensive. A troubleshooting section addresses ABI compatibility import errors, missing headers or libraries, build resource issues, verbose build logging, UV/virtual-environment problems including cuDNN sublibrary loading failures, and JAX FFI registration errors.
A breaking change is documented for v1.7: the PyTorch padding mask definition changed so that True now means masking out a position rather than including it, unifying mask semantics across frameworks. Convergence notes state that FP8 and MXFP8 showed no significant difference from BF16 training loss curves in tested configurations, with validation on downstream LLM tasks, and list models such as MPT-1.3B, Llama2-7B, LLM-8B, MPT-13B, MoE-16B and Llama2-70B across frameworks including Mosaic Composer, Alibaba Pai and Megatron Core.
Integrations listed include DeepSpeed, Hugging Face Accelerate, Lightning, MosaicML Composer, NVIDIA JAX Toolbox, NVIDIA Megatron-LM, NVIDIA NeMo Megatron Bridge, Amazon SageMaker Model Parallel Library, Levanter, GPT-NeoX and Hugging Face Nanotron. The project is Apache-2.0 licensed and welcomes contributions via its CONTRIBUTING guide.
Comments
0 people shared their preference · Deer Point appears after 10 participants
Sign in to join the discussion.