Gradient Checkpointing(グラディエントチェックポインティング)
訓練時のメモリ使用量を削減するため、順伝播の中間活性値を一部のみ保存し、逆伝播時に必要な部分を再計算する手法。計算時間とメモリのトレードオフを実現する。
Gradient Checkpointing(勾配チェックポインティング)とは
Gradient Checkpointing(Activation Checkpointing / Rematerializationとも呼ばれる)は、深層ニューラルネットワークの訓練においてGPUメモリ使用量を大幅に削減する手法です。通常、逆伝播(backward pass)では順伝播(forward pass)で計算した全ての中間活性値(activations)を保持する必要がありますが、Gradient Checkpointingではこれらの一部のみを保存し、必要に応じて再計算します。
メモリ問題の背景
Transformerベースの LLM では、各層の中間活性値がGPUメモリを大量に消費します。
| モデルサイズ | 層数 | シーケンス長 | 活性値メモリ(FP16) |
|---|---|---|---|
| 7B (LLaMA) | 32 | 2048 | 約30GB |
| 13B | 40 | 2048 | 約50GB |
| 70B | 80 | 4096 | 約200GB |
A100 80GBのGPUでは、7Bモデルでもバッチサイズ1で活性値だけでメモリの大半を消費します。Gradient Checkpointingにより、この活性値メモリを60-70%削減できます。
動作原理
チェックポイントなし(通常の訓練)
全N層の活性値をメモリに保持します。メモリ使用量はO(N)で、層数に比例します。
チェックポイントあり
√N個の層の活性値のみ保持(チェックポイント)し、その間の活性値は逆伝播時に再計算します。メモリ使用量はO(√N)に削減されますが、forward passが約33%増加します。
| 方式 | メモリ使用量 | 計算コスト | 訓練時間への影響 |
|---|---|---|---|
| チェックポイントなし | O(N) ≈ 100% | 1× forward | 基準 |
| 全層チェックポイント | O(√N) ≈ 30-40% | ~1.33× forward | +20-33% |
| 選択的チェックポイント | O(k) ≈ 50-60% | ~1.15× forward | +10-15% |
チェックポイント戦略
均一チェックポイント(デフォルト)
N層ある場合、√N層ごとにチェックポイントを配置します。メモリ削減と再計算コストのバランスが最も良い配置です。32層モデルでは約6層ごとにチェックポイント(計5-6個)が配置されます。
選択的チェックポイント
全層ではなく、メモリ消費の大きい層(Self-Attention層)のみにチェックポイントを配置します。FFN層の活性値は保持したまま、Attention層のみ再計算することで、メモリ削減効果と計算コスト増加のバランスを最適化します。
セグメントチェックポイント
モデルを固定サイズのセグメントに分割し、セグメント境界のみにチェックポイントを配置します。DeepSpeedのActivation Checkpointingで採用されている方式です。
主要フレームワークでの実装
| フレームワーク | API | 設定方法 |
|---|---|---|
| PyTorch | torch.utils.checkpoint.checkpoint | 関数ラッパーで個別指定 |
| HuggingFace | gradient_checkpointing=True | Trainerの引数1つ |
| DeepSpeed | activation_checkpointing config | JSON設定ファイル |
| Megatron-LM | --checkpoint-activations | コマンドラインフラグ |
Gradient AccumulationやFSDPとの併用
Gradient Checkpointingは他のメモリ最適化手法と組み合わせることで、さらに大きなモデルを限られたGPUで訓練可能にします。
| 手法の組み合わせ | 7Bモデルのメモリ(A100 80GB) | バッチサイズ |
|---|---|---|
| なし | 約70GB(バッチ1) | 1 |
| Checkpointing のみ | 約45GB | 2-3 |
| Checkpointing + FSDP | 約20GB/GPU(4GPU) | 4-8/GPU |
| Checkpointing + FSDP + Accumulation | 約20GB/GPU | 実効256 |
FAQ
Q1: Gradient Checkpointingによる訓練時間の増加はどの程度ですか?
全層チェックポイントで20-33%、選択的チェックポイントで10-15%の増加が一般的です。メモリ制約によりバッチサイズを上げられない場合、チェックポイントでメモリを節約してバッチサイズを増やした方が、総訓練時間は短くなることもあります。
Q2: 推論時にもGradient Checkpointingは使いますか?
いいえ。推論時は逆伝播が不要なため、中間活性値の保持自体が不要です。Gradient Checkpointingは訓練時のみ有効な最適化です。
Q3: Fine-tuningでもGradient Checkpointingは有効ですか?
はい。特にLoRAを使わないフルパラメータFine-tuningでは、活性値メモリの削減効果が大きく、消費者向けGPU(RTX 4090 24GB等)でも7Bモデルの Fine-tuning が可能になります。