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.

NVIDIA Transformer Engine

概要 – 高度に最適化されたカーネルと自動混合精度APIを提供することで、NVIDIA GPU上でのTransformer型ニューラルネットワークを高速化するライブラリです。FP8、MXFP8、NVFP4などの低精度フォーマットを使用して、大規模言語モデル(LLM)、混合エキスパート(MoE)モデル、マルチモーダルTransformerのトレーニングと実行を可能にします。これにより、精度をFP16/BF16と同等に維持しつつ、メモリ使用量の削減とスループットの向上を実現します。

主な機能

  • FP8優先サポート: Hopper、Ada、Blackwell GPUでのFP8サポートに加え、Blackwellでは新しいMXFP8/NVFP4フォーマットをサポート。
  • フレームワークに依存しないC++コア: PyTorchおよびJAX/Flax向けの軽量なPythonバインディングを提供。
  • Fused kernels: 複数の操作を単一のGPU起動に統合し、高速化を実現(例:FlashAttention-2/-3)。
  • 自動スケーリング係数処理: te.autocastを有効にするだけで、ライブラリがFP8トレーニングに必要な計算を管理。
  • 主要なLLMスタックとの統合フック: DeepSpeed、Hugging Face Accelerate、PyTorch Lightning、MosaicML Composerなど。
  • 並列化パターンへの対応: Tensor、Sequence、Context並列およびMoEワークロードのサポート。

典型的なワークフロー (PyTorchの例)

import torch, transformer_engine.pytorch as te
from transformer_engine.common import recipe

model = te.Linear(768, 3072, bias=True)
inp   = torch.randn(2048, 768, device='cuda')
fp8_recipe = recipe.DelayedScaling(margin=0, fp8_format=recipe.Format.E4M3)

with te.autocast(enabled=True, recipe=fp8_recipe):
    out = model(inp)
loss = out.sum(); loss.backward()

JAX/Flaxでも同様のパターンが使用可能で、te.autocastがフォワードパスをラップします。

インストール

  • Docker (推奨): NVIDIA NGCコンテナ (nvcr.io/nvidia/pytorch:26.01-py3 または nvcr.io/nvidia/jax:26.01-py3) をプルしてください。エンジンは /opt/transformerengine 内にプリインストールされています。
  • pip: pip install --no-build-isolation transformer_engine[pytorch] (または [jax] / 両方)。通常のCUDA、cuDNN、C++17ツールチェーンを使用してソースからのインストールが可能です。
  • conda: conda install -c conda-forge transformer-engine-torch (JAXサポートは近日公開予定)。

使用すべき場面

  • GPUメモリや計算帯域幅が制限となるLLMまたはMoEモデルのトレーニング。
  • レイテンシとメモリフットプリントが重要となる推論パイプラインのデプロイ(特にFP8ハードウェアを備えたHopper/Blackwell GPU)。
  • モデルコードを書き換えることなく、低精度へのアップグレードを即座に導入したいPyTorchまたはJAXを使用中のプロジェクト。

制限事項 / 注意点

  • FP8機能には、Compute Capability 8.9以降のGPU (Ada/Hopper/Blackwell) が必要です。
  • ソースからのビルドはメモリを大量に消費する場合があります (FlashAttention-2のコンパイルなど)。OOMが発生した場合は MAX_JOBS=1 を設定してください。
  • PyTorchとエンジンの間のABI不一致はインポートエラーの原因となる可能性があります。両方が同じC++ ABIでビルドされていることを確認してください。

リソース


Transformer Engineは、最先端の低精度ハードウェアを使用して、現代的なTransformerワークロードの加速に焦点を当てた、NVIDIAがメンテナンスしているオープンソースプロジェクトです。

関連

  • Dispatch
  • Dispatch
  • Dispatch
  • プロジェクト
  • Dispatch