FSDPとは
FSDP(Fully Sharded Data Parallel)とは、PyTorchが提供する分散学習の手法で、モデルのパラメータ、勾配、オプティマイザ状態をすべてのGPUに分散(シャーディング)することでメモリ効率を最大化する技術です。MicrosoftのDeepSpeed ZeRO Stage 3に相当する機能をPyTorchネイティブで実現します。
| ひとことで言うと | モデルの重みや勾配もGPU間で分割して持つ、省メモリな分散学習手法。 |
|---|---|
| 何がうれしいか | 各GPUのメモリ使用量が大幅に減り、より大きなモデルを学習できる。 |
| 注意点 | 必要なときに集める通信が増える。高速な接続環境が望ましい。 |
FSDPの仕組み
通常のデータ並列(DDP)では各GPUがモデル全体のコピーを保持しますが、FSDPではモデルパラメータをGPU間で分割して保持します。フォワードパスやバックワードパスで必要になった際にAllGather操作でパラメータを一時的に再構成し、計算後にはメモリから解放します。これにより、単一GPUのメモリに収まらない大規模モデルの学習が可能になります。
1通常:全GPUが全部保持重み・勾配・最適化状態
▶
2FSDP:分割して保持各GPUは一部だけ
▶
3必要時に集める計算する層だけ復元
▶
4使い終わったら解放メモリを空ける
▶
5大きなモデルが載る同じGPU数でより大規模に
FSDPの利点
FSDPの最大の利点はPyTorchネイティブであることです。外部ライブラリへの依存なしに大規模モデルの分散学習が可能で、PyTorchのエコシステム(自動微分、モジュールAPI、チェックポイント)とシームレスに統合されます。Mixed Precision学習やActivation Checkpointingとの組み合わせも容易です。
PyTorch FSDPの進化
PyTorch 2.0以降ではFSDP2が開発されており、通信の効率化やAPIの簡素化が進んでいます。Metaの大規模言語モデル(LLaMAなど)の学習にもFSDPが使用されており、実績あるスケーラブルな分散学習ソリューションとして定着しています。