JA EN
体系並列・分散
·★ 会員·13分で読めます

分散学習を1から — データ並列・モデル並列・通信がボトルネックになるとき

なぜ1台に載らないのかをメモリの内訳から出発し、データ並列とall-reduce、ZeRO/FSDPが何を分割するのか、テンソル並列とパイプライン並列を順に。最後は通信量と計算量の比で「どこで頭打ちになるか」を自分で見積もり、勾配蓄積・NCCLの設定・ハングの切り分けまで降ります。

対象textタスクtraining

比喩: 百人で一冊の辞書を書き換える

百人で1冊の辞書を改訂する仕事を考えます。分担して書くのは簡単です。難しいのは突き合わせのほうで、各自の修正を全員に反映しないと矛盾した辞書ができてしまう。1項目直すたびに全員で会議を開けば、会議の時間が執筆の時間を超えます。人数を増やすほど会議が長くなり、ある人数から先は増員が逆効果になる。

分散学習で起きるのも同じことです。計算を分けるのは難しくありません。難しいのは、分けた結果を全員のモデルに一致させて戻すところ。この記事の主役は演算器ではなく通信で、見積もるべき比は「計算時間 ÷ 通信時間」です。

なぜ1台に載らないのか

まずメモリの内訳を数えます。パラメータ数を NN とし、混合精度でAdamを使う典型的な構成を置きます。パラメータをbf16で持つと 2N2N バイト、勾配も同じ精度なら 2N2N。ここに、更新の精度を保つためのfp32マスターパラメータ 4N4N、Adamの一次モーメント 4N4N、二次モーメント 4N4N が乗ります。

Mstate2N+2N+4N+4N+4N=16N [バイト]M_{\text{state}} \approx 2N + 2N + 4N + 4N + 4N = 16N \ [\text{バイト}]
(1)

つまりこの式は、重み1個を学習させるには、重みそのもののほかに「予備の高精度版」と「過去を覚えておくメモ2枚」がくっついてくると言っています。MstateM_{\text{state}} は学習が終わるまで居座り続けるメモリ、NN は重みの個数です。

「パラメータ1個あたり16バイト」。N=7×109N = 7 \times 10^9 なら約112 GB——まだ1つも計算していないのに、状態だけで加速器1枚のメモリを超えます。しかも NN に比例するこの項とは別に、逆伝播のために取っておく中間出力(活性値)が要ります。

MactLBsdM_{\text{act}} \sim L \cdot B \cdot s \cdot d

つまり「活性値は、層を深くしても、まとめて流す量を増やしても、文を長くしても、同じ勢いで膨らむ」ということ。足し算ではなく掛け算で効くのがこの項の性格です。

LL は層数、BB はマイクロバッチ、ss は系列長、dd は隠れ次元。層数とバッチと系列長の積で効くので、モデルが同じでもバッチを倍にすれば倍になります。ここを削る定石が活性値チェックポイントで、層の入力だけ残して中間を捨て、逆伝播時に再計算する。メモリは大きく減り、計算は再計算ぶん(おおむね順伝播1回ぶん)増えます。

要するに、1台に載らない理由は2つあり、打ち手も別です。16N16N の状態は分割で減らし、活性値は再計算とバッチ調整で減らす。 この区別を最初に付けておくと、あとの選択がぶれません。

データ並列とall-reduce

いちばん素直な分割はデータのほうです。KK 台がそれぞれモデルの完全なコピーを持ち、別々のミニバッチで順伝播と逆伝播をして、勾配だけを平均して揃える。これがデータ並列で、モデルが1台に載るあいだは常に第一候補です。

揃える操作をall-reduceと呼びます。全員の値を足して(reduce)、その結果を全員が持つ(broadcast)。素朴に1台へ集めて配り直すとその1台の帯域が詰まるので、実装はリングall-reduceを使います。KK 台を輪にして、データを KK 個の塊に切り、隣へ送りながら足し込む段(reduce-scatter)と、完成した塊を隣へ回す段(all-gather)の2周。1台が送るバイト数は

Vcomm=2K1KNb2NbV_{\text{comm}} = 2\,\frac{K-1}{K}\, N b \approx 2Nb
(2)

つまり「どの台も、モデル全体の勾配をおよそ2回ぶん送り出すだけ」ということ。VcommV_{\text{comm}} は1台が1ステップで回線に流すバイト数です。

bb は勾配1要素のバイト数です。KK が大きくなっても 2Nb2Nb に漸近するだけで、台数にほぼ依存しません。 32台でも1024台でも1台が送る量は変わらない——リングall-reduceが分散学習の標準になった理由がこれです(ただし段数は KK に比例するので、レイテンシは効きます)。

通信量と計算量の比

頭打ちの位置は、1ステップの計算時間と通信時間を比べれば出ます。行列積のコストで見たとおり、学習の計算量はトークン1個あたり約 6N6N FLOPs。1台が1ステップで処理するトークン数を DD、1台の実効演算性能を PP とすると計算時間は 6ND/P6ND/P、通信時間は上の 2Nb/BW2Nb/\mathrm{BW} です。比を取ると

TcompTcomm=6ND/P2Nb/BW=3DBWbP\frac{T_{\text{comp}}}{T_{\text{comm}}} = \frac{6ND/P}{2Nb/\mathrm{BW}} = \frac{3 D \,\mathrm{BW}}{b\,P}
(3)

つまりこの分数は「1ステップの計算が、通信の何倍あるか」を表しています。1を大きく上回っていれば通信は計算の裏に隠れ、1を下回れば台数を足しても待ち時間が伸びるだけ、ということ。

NN が消えました。 モデルを大きくしても計算と通信は同じ割合で増えるので、比は変わらない。効くのは3つだけです——1台あたりのトークン数 DD相互接続の帯域 BW\mathrm{BW}、そして勾配の精度 bb

読み方は素直です。台数を増やしても全体のバッチが同じなら が減り、比が悪化して通信律速へ倒れる。逆にマイクロバッチを大きくすれば が増えて通信が隠れる。「台数を倍にしたのに1.2倍しか速くならない」の典型的な正体はこれで、[GPUの実行モデル](/ja/a/gpu-execution-model/)で見たAmdahlの直列部分に、通信がそのまま座っている状態です。

この先にあるもの

§

ここから先は会員限定です

解説記事371本・教科書26章・学生モード48単元・論文精読6本が、月額¥490ですべて読み放題になります。新しい解説は毎日3本ずつ増えます。いつでも解約でき、解約後も期間の終わりまで読めます。

会員の方はログインすると続きが表示されます

コメント

コメントにはログインが必要です