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

2026年9月20日

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
  • 主な設定パラメータ (DeepSeek-V3再現時):
    • te_moe_block: truete_gmm_quantization: "te_mxfp8"ragged_buffer_factor: 2.0sparse_matmul: trueprefuse_moe_weights: trueweight_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: 下書きに戻していた記事を、材料を集め直して書き直し、公開に戻しました。