Skip to main content
QUICK REVIEW

[論文レビュー] Robustness to Pruning Predicts Generalization in Deep Neural Networks

Lorenz Kuhn, Clare Lyle|arXiv (Cornell University)|Mar 10, 2021
Machine Learning and Algorithms参考文献 33被引用数 9
ひとこと要約

本稿では、訓練損失を増加させずに保持できるパラメータの最小割合として定義される、ニューラルネットワークの単純さを測る新しい指標「prunability」を導入する。この指標は、多様なモデルにおいて一般化性能を強く予測し、特に過パラメータ化された設定やダブルデセントの状況において、従来の複雑さの指標を上回る性能を示す。

ABSTRACT

Existing generalization measures that aim to capture a model's simplicity based on parameter counts or norms fail to explain generalization in overparameterized deep neural networks. In this paper, we introduce a new, theoretically motivated measure of a network's simplicity which we call prunability: the smallest \emph{fraction} of the network's parameters that can be kept while pruning without adversely affecting its training loss. We show that this measure is highly predictive of a model's generalization performance across a large set of convolutional networks trained on CIFAR-10, does not grow with network size unlike existing pruning-based measures, and exhibits high correlation with test set loss even in a particularly challenging double descent setting. Lastly, we show that the success of prunability cannot be explained by its relation to known complexity measures based on models' margin, flatness of minima and optimization speed, finding that our new measure is similar to -- but more predictive than -- existing flatness-based measures, and that its predictions exhibit low mutual information with those of other baselines.

研究の動機と目的

  • 過パラメータ化された深層ニューラルネットワークにおいて、パラメータノルムやパラメータ数といった従来の一般化測度が一般化を予測できないという問題に対処する。
  • モデルの単純さを理論的に裏付けられた新しい指標を特定し、より大きなモデルが一般化性能が悪くなると誤って予測するという直感に反する予測を回避する。
  • CIFAR-10およびCIFAR-100における多様な畳み込みネットワークを用いて、実験的にprunabilityがテスト性能と強い相関を示すことを示す。
  • prunabilityが、平坦性、マージン、最適化速度といった従来の指標とは異なる誘導的バイアスを捉えているかどうかを調査する。
  • 一般化行動が古典的な誘導的バイアスに従わない、挑戦的なダブルデセントの状況において、prunabilityを評価する。

提案手法

  • 訓練損失が上昇しない範囲で保持できるパラメータの最小割合としてprunabilityを定義する(例:マグニチュードベースの反復的プルーニングを用いる)。
  • 各訓練済みモデルに対して反復的マグニチュードプルーニングを用いてprunabilityを計算し、訓練損失が上昇するまで繰り返す。
  • 既存の一般化ベースライン(マージン、平坦性(ヘッセ行列の曲率を用いて)、最適化速度、ランダム摂動に対するロバストネス)とprunabilityを比較する。
  • 条件付き独立性検定および相互情報量分析(例: Kendall’s τ、調整済み R²、条件付き相互情報量)を用いて、prunabilityと他の指標との統計的関係を評価する。
  • 複数のアーサテクチャ(例:ResNets)およびデータセット(CIFAR-10、CIFAR-100)を用いた実験を実施し、モデル幅を変化させたダブルデセント設定を含む。
  • 同じ大きさの摂動に対して、プルーニングとランダムな重み摂動の影響を、訓練損失およびテスト損失の観点から比較して、機能的影響を分析する。

実験結果

リサーチクエスチョン

  • RQ1prunabilityは、過パラメータ化された深層ニューラルネットワークにおける一般化性能の信頼できる予測子として機能するか?
  • RQ2prunabilityは、モデルサイズが増加するにつれて悪化すると誤って予測する従来の複雑さの指標の失敗モードを回避するか?
  • RQ3prunabilityは、極小点の平坦性、マージン、最適化速度といった既存の一般化指標と比較して、どのように異なるか?
  • RQ4prunabilityの予測性能が、既存の指標との相関によって説明可能か、それとも独自の誘導的バイアスを捉えているか?
  • RQ5一般化誤差が訓練データサイズを超えるモデル容量に達した後も減少するダブルデセントの状況において、prunabilityはどのように振る舞うか?

主な発見

  • CIFAR-10で訓練された多様な畳み込みネットワークにおいて、prunabilityはテスト損失と強く相関しており、一般化性能の予測に非常に効果的である。
  • ノルムベースやパラメータ数の測度とは異なり、prunabilityはモデルサイズに伴って増加せず、より大きなモデルがより良い一般化性能を示すと正しく予測する。
  • ダブルデセントの状況においても、prunabilityはテスト性能の情報を的確に保持しており、既存の強力なベースラインを上回るか、同等の性能を示す。
  • prunabilityはランダム摂動に対するロバストネスと高い条件付き相互情報量を示すが、より予測力が高く、これはprunabilityがより関連性の高い単純さの側面を捉えていることを示唆する。
  • 同じ大きさの摂動に対して、プルーニングは訓練損失により大きな悪影響を与えるが、同時にテスト損失を改善する可能性がある。これは、関数的特徴に差があることを示しており、prunabilityの優れた性能の背後にある要因である可能性がある。
  • prunabilityは平坦性ベースの指標やマージンベースの指標と低い相互情報量を示しており、既存のアプローチが包含しない、独自の誘導的バイアスを捉えていることが示唆される。

より良い研究を、今すぐ始めましょう

論文の読解から最終レビューまで、研究時間を劇的に削減しましょう。

クレジットカード登録不要

このレビューはAIが作成し、人間の編集者が確認しました。