テクニカルガイド

勾配チェックポイント

勾配チェックポイント (アクティベーション チェックポイントとも呼ばれる) は、フォワード パス中にほとんどの中間アクティベーションを破棄し、バックプロパゲーション中にオンザフライで再計算する、メモリを節約するトリックです。

2分の読書最終更新日

概要

It lets you train deeper, larger networks by trading extra compute for much lower memory use.

ディープダイブ

バックプロパゲーションでは勾配を計算するためにニューラル ネットワークが必要となるため、トレーニング ニューラル ネットワークは通常、順方向パス中にすべての層のアクティベーションを保存します。深いモデルの場合、これらのアクティベーションはメモリを支配します。代わりに、勾配チェックポイント設定では、まばらなセットの「チェックポイント」レイヤーでのみアクティベーションが保存され、残りは破棄されます。 backprop は、アクティベーションが削除された領域に到達すると、そのセグメントに対してのみ前方計算を再実行して、必要なものを再生成してから続行します。チェックポイントがほぼ N の平方根ごとに配置されるため、アクティベーション用のメモリは N 次から N の平方根まで低下しますが、コンピューティングは前方パスが 1 回追加されるだけで増加します (およそ 20 ~ 30% 遅くなります)。これにより、より大きなバッチ サイズやより深いトランスフォーマーを同じ GPU に適合させることが可能になります。

技術的な洞察

この手法は、時間とメモリのトレードオフを利用します。すべてのアクティベーションを保存するのは高速ですが、メモリを大量に消費します。最新のアクセラレータを使用すると、メモリ不足によるコストと比較すると、それらの再計算は安価になります。 PyTorch (torch.utils.checkpoint) のようなフレームワークはモジュールをラップするため、前方出力は保存されますが、内部は後方実行中に再計算されます。チェックポイントの配置の選択は重要です。およそ sqrt(N) 個のセグメントを均等な間隔で配置することで、合計のメモリを最小限に抑えながら、全体のコンピューティングの前方パスを 1 つだけ追加するだけになります。

戦略的影響

費用と予算

アーキテクチャの決定により、パフォーマンスと運用コストが何年にもわたって推進されます。

より明確な判決

技術教育は、チームが最新のスタックだけでなく、適切なスタックを選択するのに役立ちます。

品質管理

より良いエンジニアリングの選択により、本番環境での信頼性に関するインシデントが減少します。

勾配チェックポインティングの将来

勾配チェックポイントは現在、大規模モデルのトレーニングの標準であり、ライブラリが最適なチェックポイントの場所を選択することで自動化が進んでいます。 FSDP、混合精度、オフロードと自然に組み合わせて、モデルのサイズを大きくします。負荷の高い操作 (アテンション マトリックスなど) をキャッシュしたままにして、負荷の低い操作のみを再計算する「選択的」チェックポイントに加え、最適な速度とメモリのバランスを保つために何を保存するか再計算するかを自動的に決定する PyTorch の torch.compile などのツールのコンパイラ駆動のアプローチを期待します。

現実世界の実装

レイヤーのアクティベーションを破棄して再計算することにより、単一の GPU 上でより大きなバッチ サイズでディ​​ープ トランスフォーマーをトレーニングします。

アクティベーション マップが GPU メモリをオーバーフローしてしまう高解像度画像上のビジョン モデルを微調整します。

フェイス トランスフォーマーをハグすると、微調整中に gradient_checkpointing=True が有効になり、10 億のパラメーター モデルに適合します。

チェックポイントと FSDP を組み合わせることで、パラメーターとアクティベーションの両方が小さく保たれ、非常に大規模な言語モデルのトレーニングが可能になります。

リスクとガードレール

1 つのベンチマークを最適化すると、より広範なシステムの弱点が隠れる可能性があります。

インフラストラクチャとメンテナンスのコストは過小評価されがちです。

システムが複雑になるにつれて、セキュリティと可観測性のギャップが拡大する可能性があります。

実装ロードマップ

1

実装前にレイテンシ、品質、コストの目標を定義します。

2

現実的な負荷とデータ条件でのベンチマーク。

3

エラー、ドリフト、ユーザーへの影響を計測器で監視します。

4

スケーリングの前に、ロールバックとインシデント対応のパスを準備します。

探検を続けましょう

Free newsletter

Get the daily AI briefing

Three verified AI stories every weekday morning, written in plain English. Free forever, no ads.

One email each weekday. Unsubscribe in one click. We never sell or share your address.

Test yourself

Take the Gradient Checkpointing quiz

Instant feedback on every answer, and a shareable certificate with a verifiable ID once you pass a course.

クイズを開始する

Support free AI education. AI Understanding is a 501(c)(3) nonprofit — no ads, no paywall, ever. Make a donation

よくある質問

What is Gradient Checkpointing?

勾配チェックポイント (アクティベーション チェックポイントとも呼ばれる) は、フォワード パス中にほとんどの中間アクティベーションを破棄し、バックプロパゲーション中にオンザフライで再計算する、メモリを節約するトリックです。追加のコンピューティングと引き換えにメモリ使用量を大幅に削減することで、より深く大規模なネットワークをトレーニングできます。

メモリを節約するために、勾配チェックポインティングは主に何を交換しますか?

勾配チェックポイントは、バックワード パス中に破棄されたアクティベーションを再計算し、メモリを大幅に削減する代わりに余分な計算を消費します。

通常、アクティベーションが転送パス中に保存されるのはなぜですか?

Backprop は、フォワード パスからの中間アクティベーションを使用して勾配を計算するため、それらは再計算されない限り利用可能である必要があります。

N 層ネットワークの sqrt(N) 層ごとにチェックポイントが配置された場合、アクティベーション メモリは大まかにどのように拡張されますか?

N の平方根ごとにチェックポイントを配置すると、格納されるアクティベーション メモリが N 次から sqrt(N) 次まで減少します。

適切に配置されたグラデーション チェックポイントにより、通常どのくらいの追加コンピューティングが追加されますか?

チェックポイントを適切に配置すると、オーバーヘッドはおよそ 1 回の追加のフォワード パスに相当し、多くの場合約 20 ~ 30% の速度が低下します。

PyTorch では、モジュールに勾配チェックポイントを適用するためにどのユーティリティが一般的に使用されますか?

torch.utils.checkpoint はモジュールをラップして、その内部アクティベーションが保存されるのではなく逆方向に再計算されるようにします。