マイクロバッチ、後方パス、更新を切り分ける
マイクロバッチとは、レプリカ上の前方パスで処理されるサンプルのグループです。後方パスはその勾配への寄与を計算します。オプティマイザの更新は、利用可能な勾配を用いてパラメータを変更します。蓄積を行う場合、この更新の前に複数の前方/後方パスが行われ、そのグループの間はパラメータは変化しません。
PyTorch では、勾配はそのために用意されたテンソルに蓄積されます。したがって、各マイクロバッチの後に勾配を消去すると、意図した蓄積が打ち消されてしまいます。逆に、2つのグループの間でゼロに戻すのを忘れると、前回の更新のサンプルが寄与してしまいます。
ログの単位を定義しましょう。マイクロバッチ番号、オプティマイザの更新、見たサンプル数、トークン数のいずれかです。「step」という言葉だけでは曖昧です。一方の実験で軸が8倍のサンプル数を表していると、損失曲線を正しく比較できません。
技術的な出典: PyTorch — 勾配の蓄積とゼロリセット
GPUを二重に数えずに実効バッチを計算する
マイクロバッチあたりかつレプリカあたりのサンプル数をm、蓄積されたマイクロバッチの数をA、データ並列のレプリカ数をDとします。これらのサイズが一定で、サンプルが正しく配分されている場合、1回のグローバル更新に寄与するサンプル数はm × A × Dです。
係数Dは必ずしもマシンのすべてのカードを指すわけではありません。テンソル並列またはパイプライン並列で同じモデルを共有するGPUは、それぞれがデータレプリカになるわけではありません。実際に構成されたグループを記載し、単にバッチの販売数量を記載しないでください。
算数の例:マイクロバッチあたり2サンプル、8回の蓄積、2つのレプリカで、グローバル更新あたり32サンプルになります。各レプリカはこのグループで16サンプルを処理します。表はカウントを比較しているだけで、そのメモリや速度をランク付けしているのではありません。
| m | A | D | グローバル更新あたりの例数 |
|---|---|---|---|
| 2 | 8 | 2 | 32 |
| 1 | 16 | 2 | 32 |
| 4 | 4 | 2 | 32 |
| 2 | 8 | 1 | 16 |
技術的な出典: PyTorch — 混合精度における実効バッチとアキュムレーション · PyTorch — DistributedDataParallelのレプリカの挙動
実際に評価された要素に基づいて損失を正規化する
同じ数の関連要素を含むマイクロバッチの平均損失では、各寄与をAで割るとグループの平均になります。このルールは、フレームワークがすでにこの正規化を行っていないことを前提としています。アキュムレーションを管理するツールを使う場合は、手動で除算を追加する前にその契約を確認してください。
トークンごとの平均損失では、長さが異なると分母が変わります。損失の合計を、グループ内で実際に教師ありのトークン(パディングと無視された位置を除く)で割る必要があります。マイクロバッチの平均の平均は、一般に同じ目的関数にはなりません。
理論上の例:あるマイクロバッチは教師ありトークン512個で平均損失が2、別のマイクロバッチは1,536個で平均が4です。加重平均は (512 × 2 + 1,536 × 4) ÷ 2,048 = 3.5 となります。非加重平均は3で、小さいグループを過大評価します。これらの値は計算の説明のみを目的としています。
完全なアキュムレーショングループを構成する
まずグループの境界とその分母を用意します。各マイクロバッチについて、出力、正規化された損失、および中間更新なしの逆伝播を計算します。不要になった出力は解放してください。損失をそのグラフに紐付けたままリストに保持すると、割り当ての寿命が延びる可能性があります。
最後の寄与の後、完全な勾配に対して予定された操作を適用し、その後更新を行います。次に、次のグループのために勾配をリセットします。走査がA個未満のマイクロバッチで終わる場合は、その部分グループを実際の分母で処理するか、除外するかを明示的に選択してください。対象となる例を記録しておきます。
GradScalerを使う混合精度では、スケール係数はアキュムレーション中は一定のままです。実際のアンスケールと必要に応じたクリッピングは寄与の後に行われ、scalerの更新はステップの試行に続きます。非有限値のチェックにより、パラメータの変更が妨げられることがあります。
スケジューラは、ループが宣言する単位に従う必要があります。オプティマイザの更新ごとに定義されている場合、各マイクロバッチで呼び出すとスケジュールが変わってしまいます。システムがステップをスキップできる場合は、試行と実際に適用された更新を別々に記録してください。
技術的な出典: PyTorch — autogradグラフとbackward用に保持されるテンソル · PyTorch — アキュムレーション、アンスケール、クリッピング、GradScaler
マルチGPUでは、勾配の縮約と例の分配を確認する
DistributedDataParallelはレプリカ間で勾配を同期します。通常の挙動では、それらを平均化します。したがって、ローカルに合計された損失とローカルに平均された損失ではスケールが異なります。レプリカごとにトークン数が異なる場合は、グローバルな分母とこの縮約を併せて考慮する必要があります。
実際に処理されるIDを確認してください。同じ例をすべてのカードで意図せず重複させても、グループの情報量はそれだけ増えません。中間通信を遅延させるには、最終同期の前のマイクロバッチでno_syncを使用できます。そのコンテキストは順伝播もカバーする必要があります。
このルールをすべての分散システムに当てはめないでください。状態のシャーディング、パイプライン、通信フック、フレームワークによって実際の操作は変わり得ます。まずツールがサポートする構成から始め、その後各レプリカの完全なグループを検証してください。
同じ実効バッチでも同じ実験が保証されない理由
等式 m × A × D は回数の数え上げです。大規模バッチの勾配を再現するには、正しく重み付けされた寄与、グループ中は同一のパラメータ状態、そしてこの分解と互換性のある演算が必要です。数値的な近さは適切な許容誤差で確認するものであり、積だけで導かれるものではありません。
BatchNormは自身のフォワードパスの入力から統計量を計算します。複数の小さなマイクロバッチでは、大規模バッチと同じグループが提示されるわけではありません。ランダム演算、計算順序、丸めも異なる場合があります。最終的な重みがビット単位で同一になると約束しないでください。
グローバルバッチの変更は、同じ学習サンプル数に対して更新回数も変え得ます。比較の軸と品質の基準をあらかじめ定めてください。学習率、スケジューラ、学習時間を同時に変更し、その新しい前提を文書化しないのは避けてください。
技術的な出典: PyTorch — BatchNorm1d の統計量 · PyTorch — 再現性の限界
メモリを確認して次に進むか判断する
バックワードパスと最初の更新を含むグループ、続いて後続のグループを計測します。フォワードの成功は勾配やオプティマイザが生成する状態を検証するものではありません。カウンタはデバイスごとの範囲を保つ必要があります。メモリの資料にはベースラインとピークの読み取り方法が示されています。
グループが失敗した場合は、マイクロバッチを縮小し、目指す実効バッチを維持するために A を再計算します(この選択が依然として妥当な場合)。この変更はピークの比例的な分割や所要時間の短縮を保証するものではありません。重みや状態が支配的な場合、蓄積だけでは不十分なことがあります。
キャンペーン前には、小規模で管理されたデータセットで更新を検証してください。同じサンプル、重み付き損失、有限の勾配、グループ境界、ステップ数です。その後、採用したプロトコルで品質を評価します。ダウンロード可能なフォルダにある推論用の小型MLPはこの学習レシピを実行せず、この検証の代わりにはなりません。
最終的な記録には m、A、D、教師ありトークン、精度、正規化、最終グループの扱い、観測されたピークをまとめます。その後、設定の選択と実験の予算に戻り、準備、試行、実際に採用された結果を分けて整理します。
- 損失が異常に小さい:蓄積による二重除算を疑う。
- 分割方法で結果が変わる:無視されたトークンと平均の平均を確認する。
- 蓄積が効かない:zero_grad と optimizer.step を確認する。
- メモリが増え続ける:マイクロバッチ間で保持された参照を探す。
- スケジュールがずれる:バックワードパスと更新を区別する。
技術的な出典: Hugging Face — 学習のメモリの目安
実用的な質問
16マイクロバッチを蓄積するとメモリも16倍になりますか?
必ずしもそうではありません。寄与は順次処理されます。勾配は保持されますが、不要なアクティベーションは解放できます。ただしピークはモデル、保持された参照、オプティマイザの状態に依存するため、グループ全体を計測してください。
2枚のGPUで常に実効バッチが2倍になりますか?
これらのGPUが、アナウンスされたマイクロバッチで2つのデータレプリカとして参加している場合に限ります。同じモデルを共有するカードが自動的に2つのレプリカになるわけではありません。分散グループと処理されたサンプルを確認してください。
すべての損失を蓄積回数で割ってよいですか?
この単純なルールは、フレームワークで既に処理されている正規化がなく、重みが等しいマイクロバッチに対応します。教師ありトークン数が変動する場合や最終グループが不完全な場合は、目的関数の実際の分母を使用してください。