JAXとNVIDIA Transformer EngineによるDropless MoE学習の高速化

基本情報
| 項目 | 内容 |
|---|---|
| 公開元 | NVIDIA Developer |
| 出典の種類 | 公開元の一次情報(NVIDIA公式ブログが一次情報) |
収集・肉付けの時点で当サイトのコードが確定させた値です。日付はJST。
概要
NVIDIAは、JAXとNVIDIA Transformer Engineを組み合わせることで、Dropless MoE(Mixture of Experts)学習を大幅に高速化できることを発表しました。この手法は、DeepSeek-V3 671Bの学習において、NVIDIA GB300環境下でスループットを10.4倍向上させ、1,024 GPU規模で97%のスケーリング効率を達成しています。
主張と根拠
発表者による測定では、JAXとTransformer Engineを用いた最適化により、DeepSeek-V3の学習パフォーマンスが大幅に向上したと報告されています。従来の最適化されていないJAXベースのスタックでは、NVIDIA GB300上での性能は103 TFLOPS/GPUに留まり、GPU間の通信が累積カーネル時間の84%を占めていました。しかし、Transformer Engineによるターゲットカーネル最適化を適用することで、1,068 TFLOPS/GPUへと、10.4倍のスループット向上を実現しています。
主な最適化技術とその根拠は以下の通りです。
- Grouped GEMMによるRagged Tensorへの対応: MoEでは各エキスパートに割り当てられるトークン数が動的に変化するため、形状が不揃いな「Ragged Tensor」が発生します。Transformer Engineの
grouped_gemm/ragged_dotは、これらを単一のカーネルコールで処理し、cuBLAS/cuBLASLtを活用してTensor Coreの利用率を最大限に高めます。 - NCCL EPによるDispatchとCombineの統合: 専門家並列(Expert Parallelism: EP)におけるトークンの転送(Dispatch)と結果の回収(Combine)を、通信バックエンドである
NCCL EPを用いて密に融合したカーネルパスとして実装しています。これにより、通信と計算のオーバーラップを図るとともに、トークンの重複排除(deduplication)によってネットワーク帯域を節約します。 - その他の最適化: JAXのホストオフローディングによるメモリボトルネックの解消や、XLAのマルチストリーム・コレクティブ(multistreaming collectives)による、NVLinkとInfiniBand間の通信のオーバーラップが挙げられます。
スケーリング性能に関しては、NVIDIA GB300 NVL72ハードウェアを用いたDeepSeek-V3 671Bの学習において、1,024 GPUを使用した場合でも97%のスケーリング効率を維持できることが示されています。
前提条件
本件の最適化結果および再現手順が成り立つ条件は、資料に基づき以下の通りです。
- 対象モデル: DeepSeek-V3 671B
- 対象ハードウェア: NVIDIA GB300, NVIDIA GB300 NVL72
- ソフトウェア・コンテナ:
- NVIDIA NGC MaxText container (
ghcr.io/nvidia/jax:maxtext-2026-09-09以降) – JAX – NVIDIA Transformer Engine – NCCL EP – XLA
- NVIDIA NGC MaxText container (
- 主な設定パラメータ (DeepSeek-V3再現時):
te_moe_block: true–te_gmm_quantization: "te_mxfp8"–ragged_buffer_factor: 2.0–sparse_matmul: true–prefuse_moe_weights: true–weight_dtype: "bfloat16"–mu_dtype: "bfloat16"
手元で再現できる範囲
読者は、NVIDIAが提供するNGC MaxTextコンテナを使用することで、最適化されたJAX MoEパスを自身の環境で再現・構築することが可能です。
基本的な使用方法
MaxTextのYAML設定ファイル、またはトレーニングスクリプトのコマンドライン引数として、以下のフラグを追加することで TE MoEBlock を有効化できます。
te_moe_block: true
te_gmm_quantization: "te_mxfp8"
ragged_buffer_factor: 2.0
te_ep_overflow_check_every_n_steps: 20
sparse_matmul: true
prefuse_moe_weights: true
DeepSeek-V3 671B の性能再現
資料では、DeepSeek-V3 671Bのベンチマーク結果を正確に再現するための詳細な設定が公開されています。これにはMaxTextの設定に加え、XLAフラグおよび環境変数の指定が含まれます。
MaxText設定例 (抜粋):
model_name: "deepseek3-671b"
max_target_length: 4096
hardware: "gpu_multiprocess"
per_device_batch_size: 6
gradient_accumulation_steps: 1
steps: 15
attention: "cudnn_flash_te"
remat_policy: "custom"
quantization: "te_fp8_currentscaling"
te_moe_block: true
te_gmm_quantization: "te_mxfp8"
ragged_buffer_factor: 2.0
te_ep_overflow_check_every_n_steps: 20
prefuse_moe_weights: true
weight_dtype: "bfloat16"
mu_dtype: "bfloat16"
#... (その他の詳細なパラメータ)
XLAフラグ設定例:
xla_gpu_all_reduce_combine_threshold_bytes: 33554432
xla_gpu_all_gather_combine_threshold_bytes: 6442450944
xla_gpu_reduce_scatter_combine_threshold_bytes: 201326592
xla_gpu_experimental_enable_nccl_symmetric_buffers: false
xla_gpu_enable_command_buffer: "'FUSION,CUBLAS,CUDNN,DYNAMIC_SLICE_FUSION'"
xla_gpu_experimental_max_unroll_factor: 8
xla_gpu_memory_limit_slop_factor: 99
環境変数設定例:
XLA_PYTHON_CLIENT_MEM_FRACTION: 0.88
CUDA_DEVICE_MAX_CONNECTIONS: 16
XLA_PJRT_GPU_HOST_MEMORY_PREALLOCATE: false
XLA_PJRT_GPU_HOST_MEMORY_LIMIT_GB: 180
ただし、これらの最適化はNVIDIA Blackwellアーキテクチャ(GB300等)に特化したものであり、他のハードウェア環境における動作については明示されていません。
資料が触れていないこと
- JAX以外の主要なディープラーニングフレームワーク(PyTorch等)を用いた場合の、同様のDropless MoE最適化との比較。
- NVIDIA製GPU以外のハードウェア環境におけるパフォーマンス。
- 学習プロセス全体にかかる具体的な計算コストや、インフラ構築のコスト。
- 従来のcapacity-based MoE手法を、同じNVIDIA GB300環境で実行した場合の具体的なパフォーマンス数値。
関連記事
出典
更新履歴
- 2026-09-20: 下書きに戻していた記事を、材料を集め直して書き直し、公開に戻しました。

