勾配累積(コウバイルイセキ)
大きなバッチサイズを直接処理するメモリが不足する場合に、バッチを複数のマイクロバッチに分割して勾配を累積し、一定ステップ後にまとめてパラメータ更新を行う手法。
勾配累積とは
勾配累積(Gradient Accumulation)は、GPUメモリの制約下で実効的に大きなバッチサイズを実現するためのシンプルかつ効果的な手法である。LLMの訓練では大きなバッチサイズ(数千〜数万サンプル)が収束の安定性と最終精度に寄与することが知られているが、単一GPUに大バッチ全体を載せることはメモリ的に不可能な場合が多い。
勾配累積では、バッチをK個のマイクロバッチに分割し、各マイクロバッチで順伝播・逆伝播を実行して勾配を加算(累積)する。K回の累積後に初めてオプティマイザのステップ(パラメータ更新)を行う。数学的にはバッチサイズB = K × マイクロバッチサイズと等価になる。
動作の流れ
通常の訓練(バッチサイズ32)
- 32サンプルを一度にGPUに載せる
- 順伝播 → 逆伝播 → 勾配計算
- オプティマイザで重み更新
- 勾配をゼロクリア
勾配累積(マイクロバッチ8 × 4累積 = 実効32)
- 8サンプルを載せて順伝播 → 逆伝播 → 勾配を累積
- 次の8サンプルで同様に勾配を累積
- さらに8サンプルで累積
- 最後の8サンプルで累積(合計32サンプル分の勾配が蓄積)
- オプティマイザで重み更新
- 勾配をゼロクリア
メモリ使用量は8サンプル分(通常の1/4)で、32サンプルと同等の訓練効果を得られる。
累積ステップ数の決め方
| 累積ステップ | 実効バッチサイズ | メモリ削減 | 訓練時間 | 適用シーン |
|---|---|---|---|---|
| 1(累積なし) | マイクロバッチサイズ | なし | 最短 | メモリに余裕がある場合 |
| 2-4 | 中程度 | 50-75% | やや増加 | 一般的なファインチューニング |
| 8-16 | 大バッチ | 87-94% | 増加 | 事前学習、大規模データ |
| 32-128 | 超大バッチ | 97%以上 | 大幅増加 | BERT/GPT事前学習 |
訓練時間の増加は理論上ゼロ(計算量は同じ)だが、実際にはGPUの利用効率低下やメモリ転送のオーバーヘッドで5〜10%程度増加する。
注意点
BatchNormalizationとの相性
BatchNormalization(BN)はマイクロバッチ単位で統計量(平均・分散)を計算する。勾配累積でマイクロバッチサイズが小さくなると、BNの統計量が不安定になり訓練品質が低下する。LLMで標準的なLayerNormやRMSNormはサンプル独立で計算するため、この問題は発生しない。
学習率のスケーリング
勾配累積で実効バッチサイズを大きくする場合、Linear Scaling Rule(バッチサイズに比例して学習率を増加)やWarmupの適用が必要。ただし累積ステップ数が2〜4程度であれば学習率の調整なしでも安定する場合が多い。
ロス関数の正規化
勾配累積時のロス計算では、各マイクロバッチのロスをマイクロバッチサイズで正規化する必要がある。PyTorchのCrossEntropyLossはデフォルトで reduction='mean' が設定されており自動正規化されるが、カスタムロス関数では手動での正規化が必要。
データ並列との違い
| 項目 | 勾配累積 | データ並列(DDP) |
|---|---|---|
| GPU数 | 1台 | 複数台 |
| 通信コスト | なし | All-Reduce |
| スループット | 逐次(遅い) | 並列(速い) |
| 実装コスト | 低い | 中程度 |
| メモリ効率 | ◎ | △(各GPUにモデル複製) |
勾配累積とデータ並列は直交する概念であり、併用可能。8GPU × 勾配累積4ステップ × マイクロバッチ4 = 実効バッチサイズ128 のように使う。
FAQ
Q1: 勾配累積は訓練精度に影響しますか?
数学的にはバッチ全体を一度に処理した場合と等価であり、精度への影響はない。ただしBatchNormalizationを使用するモデル(ResNet等)ではマイクロバッチサイズが小さくなることで統計量が不安定になり、若干の精度低下が起きうる。LLMではLayerNormを使用するためこの問題はない。
Q2: 勾配累積とDeepSpeed ZeROは併用できますか?
併用可能で、推奨される組み合わせである。ZeROがパラメータ・勾配・オプティマイザのメモリを分割し、勾配累積が活性値のメモリを削減する。DeepSpeedの設定で gradient_accumulation_steps を指定するだけで有効化できる。
Q3: 推論時にも勾配累積の概念はありますか?
推論では逆伝播がないため勾配累積は不要。ただし大量の入力を処理する際にバッチを分割して逐次処理する「推論バッチング」は概念的に類似しており、メモリ管理の手法としてContinuous Batchingなどが用いられる。