NVIDIA、Transformer Engineを活用した生物系MoEモデルの効率的学習手法を公開

NVIDIA、Transformer Engineを活用した生物系MoEモデルの効率的学習手法を公開

基本情報

項目 内容
公開元 NVIDIA Developer
公開日 2026-09-25
出典の種類 公開元の一次情報(NVIDIA公式開発者ブログであり一次情報)

収集・肉付けの時点で当サイトのコードが確定させた値です。日付はJST。

概要

NVIDIAは、生物学系基盤モデル(Biological Foundation Models)のための効率的なMixture-of-experts (MoE) 学習手法を公開しました。NVIDIA Transformer Engine (TE) を活用し、GroupedLinearによるカーネル起動オーバーヘッドの削減、MXFP8によるメモリ使用量の抑制、および複数の演算を統合するカーネル融合技術を用いることで、MoEモデルの学習効率を向上させています。

主張と根拠

NVIDIAは、TEの最適化プリミティブを用いることで、MoEアーキテクチャにおける計算の断片化、通信オーバーヘッド、およびメモリ使用量の課題を解決できると主張しています。具体的な技術的アプローチと根拠は以下の通りです。

  • GroupedLinearによる計算の効率化: Hugging Faceのベースライン実装では、Pythonのループを用いて各エキスパートに対して個別にカーネルを起動するため、オーバーヘッドが発生します。これに対し、TEのGroupedLinearは、複数のエキスパートのGEMM(行列乗算)を一つのグループ化された操作として一度に投入することで、起動およびスケジューリングのオーバーヘッドを削減します。
  • MXFP8によるメモリ削減とハードウェア加速: BF16と比較して、MXFP8は重みとアクティベーションを8ビットで表現するため、メモリ使用量を削減できます。MXFP8は32個の連続する値のブロックごとにスケーリング係数を割り当てるブロックスケーリングを採用しており、NVIDIA Blackwell GPUでは専用のTensor Core命令によってハードウェア加速されます。
  • カーネル融合によるオーバーヘッド削減: TEのSequential APIを使用すると、GroupedLinear、ScaledSwiGLU、およびルーティング重みのスケーリングを、ForwardGroupedMLP_CuTeGEMMSwiGLU_MXFP8という一つの融合カーネルに統合できます。これにより、中間データの生成を回避し、量子化やデ量子化に伴うフレームワークのオーバーヘッドを低減します。

これらの最適化を含むBioNeMoのレシピを用いた学習ベンチマークにおいて、発表者の測定では、8基のNVIDIA B200 Tensor Core GPUを使用した際、Hugging Faceのベースラインと比較して最大2.21倍のスループットを達成したと報告されています。

前提条件

本資料で示されている手法や結果には、以下の条件が含まれます。

  • 対象モデル: Mixtral-8x7B
  • 対象ハードウェア: NVIDIA B200 Tensor Core GPU
  • ソフトウェア/ライブラリ: NVIDIA Transformer Engine (TE), PyTorch, BioNeMo recipe
  • 精度設定: MXFP8 (NVIDIA Blackwell GPUを使用する場合、融合されたMXFP8 GroupedMLPカーネルを利用可能)
  • 並列化設定: Expert Parallelism (EP) を使用。8基のGPUを使用する場合、EP=8の設定などが示されています。

手元で再現できる範囲

読者は、以下のリソースや手順を通じて、MoEモデルの学習環境を構築・試行できる可能性があります。

  • コードとレシピ: NVIDIA BioNeMo Recipesに含まれる「Mixtral Native Transformer Engine recipe」を利用して、MoEベースの生物学系基盤モデルの学習を試行することが推奨されています。
  • 環境構築: NVIDIA CUDA対応環境が必要であり、提供されているDockerfileを使用するか、レシピの要件をインストールすることで環境を構築できます。
  • 実行コマンド: 以下のコマンドを用いて、環境の検証や学習を実行できることが示されています。

エキスパート並列化と学習環境が正しく動作することを確認するための、2基のGPUを用いた設定:

torchrun --nproc_per_node=2 train_fsdp2_ep.py --config-name L0_sanity

8基のGPUを用い、エキスパート並列化 (EP=8) および MXFP8 精度を適用した Mixtral-8x7B の設定:

torchrun --nproc_per_node=8 train_fsdp2_ep.py --config-name L1_8x7B_ep checkpoint.ckpt_dir=/path/to/ckpt

なお、MXFP8の融合カーネルを利用するには、NVIDIA Blackwell GPUが必要であると明記されています。

関連記事

次に読むなら

出典