AI・機械学習
上級

Gradient Checkpointingグラディエントチェックポインティング

訓練時のメモリ使用量を削減するため、順伝播の中間活性値を一部のみ保存し、逆伝播時に必要な部分を再計算する手法。計算時間とメモリのトレードオフを実現する。

0 回閲覧
0 いいね

Gradient Checkpointing(勾配チェックポインティング)とは

Gradient Checkpointing(Activation Checkpointing / Rematerializationとも呼ばれる)は、深層ニューラルネットワークの訓練においてGPUメモリ使用量を大幅に削減する手法です。通常、逆伝播(backward pass)では順伝播(forward pass)で計算した全ての中間活性値(activations)を保持する必要がありますが、Gradient Checkpointingではこれらの一部のみを保存し、必要に応じて再計算します。

メモリ問題の背景

Transformerベースの LLM では、各層の中間活性値がGPUメモリを大量に消費します。

モデルサイズ層数シーケンス長活性値メモリ(FP16)
7B (LLaMA)322048約30GB
13B402048約50GB
70B804096約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設定方法
PyTorchtorch.utils.checkpoint.checkpoint関数ラッパーで個別指定
HuggingFacegradient_checkpointing=TrueTrainerの引数1つ
DeepSpeedactivation_checkpointing configJSON設定ファイル
Megatron-LM--checkpoint-activationsコマンドラインフラグ

Gradient AccumulationやFSDPとの併用

Gradient Checkpointingは他のメモリ最適化手法と組み合わせることで、さらに大きなモデルを限られたGPUで訓練可能にします。

手法の組み合わせ7Bモデルのメモリ(A100 80GB)バッチサイズ
なし約70GB(バッチ1)1
Checkpointing のみ約45GB2-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 が可能になります。