NVIDIA/TransformerEngine
A library for accelerating Transformer models on NVIDIA GPUs, including using 8-bit and 4-bit floating point (FP8 and FP4) precision on Hopper, Ada and Blackwell GPUs, to provide better performance with lower memory utilization in both training and inference.
What it solves
Transformer Engine (TE) 解决了将 Transformer 模型(如 LLM、MoE 和多模态模型)扩展到数千亿参数时带来的高内存和计算需求。它通过利用低精度数值格式,在训练和推理过程中均能降低内存利用率并提高吞吐量。
How it works
TE 提供了一套针对 Transformer 架构高度优化的构建模块和融合内核(fused kernels)。它实现了一个自动混合精度 API,允许用户将低精度格式无缝集成到其现有的框架代码中。
关键技术能力包括:
- Low-Precision Support: 在 Hopper, Ada, 和 Blackwell GPU 上原生支持 8 位浮点数 (FP8) ,并在 Blackwell GPU 上通过 MXFP8 和 NVFP4 进一步提升效率。
- Framework Integration: 提供 PyTorch 和 JAX (Flax) 的 Python API,以及用于与其他深度学习库集成的框架无关的 C++ API。
- Automated Scaling: 内部管理 FP8 训练所需的缩放因子 (scaling factors),简化了用户的混合精度过程。
- Optimized Kernels: 包含针对 Mixture-of-Experts (MoE) 和各种并行策略(tensor, sequence, 和 context)的进阶功能及融合操作与优化。
Who it’s for
它是为使用 NVIDIA GPU(Ampere 架构或更高版本)并希望最大化硬件效率和训练速度的大规模 Transformer 模型研究人员和工程师设计的。
Highlights
- Broad Precision Support: 在最新的 NVIDIA NVIDIA GPU 上支持 FP8, MXFP8, 和 NVFP4,并在 Ampere 及更高版本上支持 FP16/BF16。
- Easy Integration: 提供用于构建具有 FP8 支持的 Transformer 层级的模块,开箱即用。 。
- C++ API: 包含一个用于深度学习库开发者的框架无关的 C++ 库。
- Performance Optimizations: 具備融合内核和对 PyTorch 中 FlashAttention-2 和 FlashAttention-3 的支持。