Skip to main content
QUICK REVIEW

[論文レビュー] SWAD: Domain Generalization by Seeking Flat Minima

Junbum Cha, Sanghyuk Chun|arXiv (Cornell University)|Feb 17, 2021
Domain Adaptation and Few-Shot Learning参考文献 60被引用数 4
ひとこと要約

本稿では、損失関数の平坦な極小値を探索することでモデルの頑健性を向上させる、新しいドメイン一般化手法SWADを提案する。密度的で過学習に注意を払った確率的重み平均化を導入することで、5つの主要なDGベンチマークで最先端の性能を達成し、従来のSOTA手法に比べて平均的なドメイン外精度を1.6パーセンテージポイント向上させた。

ABSTRACT

Domain generalization (DG) methods aim to achieve generalizability to an unseen target domain by using only training data from the source domains. Although a variety of DG methods have been proposed, a recent study shows that under a fair evaluation protocol, called DomainBed, the simple empirical risk minimization (ERM) approach works comparable to or even outperforms previous methods. Unfortunately, simply solving ERM on a complex, non-convex loss function can easily lead to sub-optimal generalizability by seeking sharp minima. In this paper, we theoretically show that finding flat minima results in a smaller domain generalization gap. We also propose a simple yet effective method, named Stochastic Weight Averaging Densely (SWAD), to find flat minima. SWAD finds flatter minima and suffers less from overfitting than does the vanilla SWA by a dense and overfit-aware stochastic weight sampling strategy. SWAD shows state-of-the-art performances on five DG benchmarks, namely PACS, VLCS, OfficeHome, TerraIncognita, and DomainNet, with consistent and large margins of +1.6% averagely on out-of-domain accuracy. We also compare SWAD with conventional generalization methods, such as data augmentation and consistency regularization methods, to verify that the remarkable performance improvements are originated from by seeking flat minima, not from better in-domain generalizability. Last but not least, SWAD is readily adaptable to existing DG methods without modification; the combination of SWAD and an existing DG method further improves DG performances. Source code is available at https://github.com/khanrc/swad.

研究の動機と目的

  • 訓練データとテストデータの分布が著しく異なるドメインシフトの課題に対処すること。
  • 通常の経験的リスク最小化(ERM)の限界を克服すること。ERMはしばしば鋭い極小値に収束し、ドメインシフト下で一般化性能が著しく低下する。
  • 理論的および実験的に、平坦な極小値がドメイン一般化(DG)の状況下でより良い一般化をもたらすことを示すこと。
  • アーキテクチャの変更なしに平坦性と一般化性能を向上させる、シンプルで効果的な手法「確率的重み平均化を密度的に拡張したSWAD(Stochastic Weight Averaging Densely)」を開発すること。
  • データ拡張や一貫性正則化手法と比較することで、性能向上の要因が平坦性に起因するものであることを確認し、ドメイン内一般化の向上ではないことを検証すること。

提案手法

  • パrameter空間の近傍における最悪ケースの経験的リスクを用いて、ドメイン一般化ギャップを上界で制約する、頑健なリスク最小化(RRM)の定式化を提案する。
  • 理論的に、最適解の近傍における最悪ケースリスクを制約することで、平坦な極小値がより小さいドメイン一般化ギャップをもたらすことを示した。
  • 確率的重み平均化(SWA)を改良し、各訓練イテレーションで重みを密度的にサンプリングすることで、損失関数の平坦な領域をよりよく探索できるようにした。
  • バリデーション損失を用いて、過学習を避けるために平均化の開始および終了イテレーションを過学習に注意を払って決定する戦略を導入した。
  • アーキテクチャの変更なしに既存のDG手法にSWADをプラグインとして適用可能であり、一貫した性能向上を実現した。
  • CPUメモリを活用して中間の重みを保存することで、GPUメモリの増加を最小限に抑えながらも、訓練効率を維持した。

実験結果

リサーチクエスチョン

  • RQ1非凸的で複雑な深層学習の損失関数の領域において、平坦な極小値を探索することでドメイン一般化ギャップを顕著に小さくできるか?
  • RQ2分布シフトがi.i.d.設定よりも顕著に強いドメインシフト下でも、平坦な極小値による一般化性能の向上が有効に機能するか?
  • RQ3アーキテクチャの変更なしに、SOTAの複雑なタスク特化型DG手法を凌駕できるシンプルで平坦性に注意を払った最適化手法(例:SWAD)は存在するか?
  • RQ4SWADの性能向上要因は、ドメイン内一般化の向上にあるのか、それともドメイン外の頑健性の向上にあるのか?
  • RQ5ドメイン一般化の文脈において、SWADは他の平坦性に注意を払った手法(例:SAM、SWA)やドメイン内一般化技術(例:Mixup、CutMix)と比較してどのように評価されるか?

主な発見

  • SWADはPACS(+2.6pp)、VLCS(+1.6pp)、OfficeHome(+4.1pp)、TerraIncognita(+3.9pp)、DomainNet(+5.6pp)の5つの主要なDGベンチマークで最先端の性能を達成し、ERMに比べて平均して1.6パーセンテージポイントのドメイン外精度向上を達成した。
  • 平均的に、SWADは既存の最良のSOTA手法に比べてドメイン外精度を1.6パーセンテージポイント向上させ、ERMベースラインに比べて3.6パーセンテージポイントの向上を達成した。
  • SWADを以前のSOTA手法(SOTA [31])と組み合わせることでさらなる向上が得られ、平均精度67.3%を達成した。これはSWAD単体の結果より0.4パーセンテージポイント高い。
  • 損失関数の可視化と平坦性指標により、SWADは通常のSWAよりも一貫して平坦な極小値を探索していることが確認された。
  • データ拡張や一貫性正則化手法(例:Mixup、CutMix)はドメイン外一般化を向上させなかったが、平坦性に注意を払った手法(SWA、SAM)は向上させた。これにより、平坦性が主因であることが裏付けられた。
  • SWADは実行時間のオーバーヘッドがERMの1.07倍から1.27倍に留まり、追加のGPUメモリコストが一切ないため、実世界の展開において実用的である。

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

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

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

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