[論文レビュー] Surrogate Gap Minimization Improves Sharpness-Aware Training
本稿では、ヘッセ行列に基づく鋭さの代わりに計算的に効率的な代理指標(サーヴィエントギャップ)を導入することで、Sharpness-Aware Minimization (SAM) の性能を向上させる新規なトレーニング手法 GSAM を提案する。この手法により、損失関数と鋭さを同時に直接最小化できる。GSAM は1回の更新で2段階の処理を実行する:(1) SAM と同様に摂動付き損失の最小化、(2) サーヴィエントギャップを低減するための直交方向への上昇。これにより、平坦な最小値が得られ、一般化性能が向上し、ImageNet における ViT-B/32 では AdamW よりトップ1正解率が 5.4% 向上する。
The recently proposed Sharpness-Aware Minimization (SAM) improves generalization by minimizing a extit{perturbed loss} defined as the maximum loss within a neighborhood in the parameter space. However, we show that both sharp and flat minima can have a low perturbed loss, implying that SAM does not always prefer flat minima. Instead, we define a extit{surrogate gap}, a measure equivalent to the dominant eigenvalue of Hessian at a local minimum when the radius of the neighborhood (to derive the perturbed loss) is small. The surrogate gap is easy to compute and feasible for direct minimization during training. Based on the above observations, we propose Surrogate extbf{G}ap Guided extbf{S}harpness- extbf{A}ware extbf{M}inimization (GSAM), a novel improvement over SAM with negligible computation overhead. Conceptually, GSAM consists of two steps: 1) a gradient descent like SAM to minimize the perturbed loss, and 2) an extit{ascent} step in the extit{orthogonal} direction (after gradient decomposition) to minimize the surrogate gap and yet not affect the perturbed loss. GSAM seeks a region with both small loss (by step 1) and low sharpness (by step 2), giving rise to a model with high generalization capabilities. Theoretically, we show the convergence of GSAM and provably better generalization than SAM. Empirically, GSAM consistently improves generalization (e.g., +3.2\% over SAM and +5.4\% over AdamW on ImageNet top-1 accuracy for ViT-B/32). Code is released at \url{ https://sites.google.com/view/gsam-iclr22/home}.
研究の動機と目的
- SAM が鋭い最小値と平坦な最小値の両方で低い摂動付き損失を示すという限界を是正するため、鋭さのより信頼性の高い指標を同定すること。
- 特に小さな近傍半径の場合に、局所最小値におけるヘッセ行列の主要固有値を近似できる、計算的に取り扱いやすいサーヴィエントギャップを定義すること。
- 摂動付き損失とサーヴィエントギャップの両方を同時に最小化するトレーニング手法を構築し、より平坦な最小値とより良い一般化性能を達成すること。
- 提案された GSAM アルゴリズムの収束を証明し、標準的な仮定のもとで SAM よりも一般化性能に優れた理論的優位性を確立すること。
- さまざまなアーキテクチャ(ResNets、ビジョントランスフォーマー、MLP-Mixer)において、GSAM の実証的妥当性を検証し、SAM やベースライン最適化手法と比較して一貫した改善を示すこと。
提案手法
- サーヴィエントギャップを $ h(w) = f_p(w) - f(w) $ として定義する。ここで $ f_p(w) $ は摂動付き損失、$ f(w) $ は標準損失であり、局所最小値におけるヘッセ行列の主要固有値を近似する。
- 勾配 $ \nabla f(w) $ を $ \nabla f_p(w) $ に平行および直交する成分に分解することで、直交方向への的確な最適化が可能になる。
- 摂動付き損失 $ f_p(w) $ における勾配降下ステップを実行し、SAM の更新ルールを保持する。
- 直交成分 $ \nabla_\perp f(w) $ 沿いの上昇ステップを適用してサーヴィエントギャップ $ h(w) $ を最小化する。直交性のおかげで $ f_p(w) $ は変化しない。
- 上昇ステップの大きさを制御するハイパーパramータ $ \alpha $ を導入し、損失と鋭さの最小化の間で柔軟なトレードオフを実現する。
- 近傍半径 $ \rho_t $ が減少するという仮定のもとで理論的収束を保証し、GSAM の収束解析を支援する。
実験結果
リサーチクエスチョン
- RQ1SAM の摂動付き損失は、鋭い最小値と平坦な最小値の区別を信頼性を持って行えるか、それとも両方で低くなることがあるか?
- RQ2トレーニング中に使用可能な、ヘッセ行列に基づく鋭さ測定の計算的に効率的な代替指標は存在するか?
- RQ3摂動付き損失と標準損失の差として定義されるサーヴィエントギャップを最小化することで、最適化が平坦な最小値へ効果的に誘導できるか?
- RQ4摂動付き損失の降下と、サーヴィエントギャップを低減する直交方向への上昇を組み合わせた2段階更新戦略は、SAM よりも優れた一般化性能をもたらすか?
- RQ5標準的な仮定のもとで、GSAM が収束可能であり、理論的に SAM よりも一般化性能に優れていることが示せるか?
主な発見
- 近傍半径 $ \rho $ が小さいとき、サーヴィエントギャップ $ h(w) = f_p(w) - f(w) $ は局所最小値においてヘッセ行列の主要固有値と理論的に同等であり、鋭さの有効かつ効率的な代理指標であることが示される。
- GSAM は ImageNet における ViT-B/32 を用いて、SAM よりもトップ1正解率が 3.2% 向上し、AdamW よりも 5.4% 向上する。これは一貫した一般化性能の向上を示している。
- アブレーションスタディにより、GSAM の上昇ステップが性能向上の主因であることが明らかになった。これは、$ \rho_t $ を一定値または減少スケジュールで設定した場合に共通して観察された。
- 実証的検証では、$ \cos\theta_t $($ \nabla f(w) $ と $ \nabla f_p(w) $ のなす角の余弦)がトレーニング全体を通して 0.9 を上回っていることが確認され、高次元パrameter空間における勾配がほぼ一致しているという理論的仮定を支持する。
- サーヴィエントギャップは $ \alpha $ が増加するにつれて減少し、訓練ステップが進むにつれて増加する傾向にある。これは、GSAM が時間経過とともにモデルをより平坦な最小値へと効果的に誘導していることを示している。
- GSAM は広範な適用性を持ち、SAM と比較して計算オーバーヘッドがほとんどないため、大規模なトレーニングにおいて実用的である。
より良い研究を、今すぐ始めましょう
論文の読解から最終レビューまで、研究時間を劇的に削減しましょう。
クレジットカード登録不要
このレビューはAIが作成し、人間の編集者が確認しました。