AI

DeepSeek-V3学習のGPU演算性能、最適化前の10.4倍に

この記事のポイント

  1. 実現したこと

    JAX上のDropless MoE学習で、割り当てが偏っても全トークンを選択先のExpertで処理する。

  2. 実現の仕組み

    Grouped GEMMが、Expertごとに異なる実トークン数の行列積を一度のカーネル呼び出しで処理する。

  3. 得られた結果

    NVIDIAのDeepSeek-V3測定では、GPU当たりの演算性能が103から1068 TFLOPSへ上がった。

  4. 従来との違い

    固定容量に合わせたトークンの破棄や余白埋めを避け、Expertごとの実トークン数を計算に使う。

Expertごとの異なるトークン数をGrouped GEMMで処理し、GPU間で送受信するMoE学習の図解。
AI生成画像

Expertごとに届くトークン数が変わるMoEモデルの学習を、JAX向けのGrouped GEMMなどで高速化した。NVIDIAのDeepSeek-V3測定では、GPU当たりの演算性能が最適化前の10.4倍となった。

Expertへの割り当てが計算量を変える

Mixture of Experts(MoE)では、ルーターがトークンごとに処理を担当するExpertを選ぶ。割り当て数はExpertごとに異なり、学習中にも変わるため、すべてのExpertを同じ大きさの行列として処理しにくい。

Dropless MoEは、割り当てが偏っても選ばれたExpertで全トークンを処理する。Expertの容量を固定する方式では、容量を超えたトークンを破棄するか、固定形状に合わせて余白を埋めることになる。

Grouped GEMMが実トークン数だけを計算する

NVIDIAのTransformer Engineは、Expertごとに長さの異なる入力をGrouped GEMMで扱う。複数のExpertの行列積を一度のカーネル呼び出しにまとめ、それぞれに実際に割り当てられたトークンの領域だけを計算する。

JAX向けの構成には、Expertの行列積に対応するMXFP8のGrouped GEMMと、グループ単位の量子化も含まれる。入力の長さが揃わないMoEの計算を、固定容量に合わせた余白埋めに頼らず進める構成だ。

GPU間の送信と回収もExpert並列処理に組み込む

Expertが複数のGPUに分かれると、ルーターの選択だけでは計算は始まらない。トークンを担当するExpertのGPUへ送り、計算された出力を回収して元のトークン順へ戻す必要がある。

Transformer Engineは、この送信と回収に最適化したExpert並列処理を提供する。Grouped GEMMが各Expert内の行列積を担い、その前後でGPU間のトークン移動を扱う。

DeepSeek-V3の測定でGPU当たり10.4倍

NVIDIAのDeepSeek-V3学習測定では、最適化前のGPU当たりの演算性能は103 TFLOPSだった。この構成では、GPU間通信が累積カーネル時間の84%を占めていた。

JAXとTransformer Engineによる最適化後、GPU当たりの演算性能は1068 TFLOPSとなり、最適化前の10.4倍に達した。結果は、トークン数が揃わないExpertの計算とGPU間の受け渡しを合わせて改善した構成の測定値である。