Skip to main content
QUICK REVIEW

[논문 리뷰] SWAD: Domain Generalization by Seeking Flat Minima

Junbum Cha, Sanghyuk Chun|arXiv (Cornell University)|2021. 02. 17.
Domain Adaptation and Few-Shot Learning참고 문헌 60인용 수 4
한 줄 요약

이 논문은 손실 곡면에서 평탄한 최소값을 찾는 방식으로 모델의 강건성을 향상시키는 새로운 도메인 일반화 방법인 SWAD를 제안한다. 조밀하고 과적합을 고려한 확률적 가중치 평균화를 도입함으로써, 다섯 개인 주요 DG 벤치마크에서 최신 기준 성능(SOTA)을 달성하며, 기존 SOTA 방법 대비 평균 도메인 외 정확도를 1.6%p 향상시킨다.

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)의 한계를 극복하기 위해, 이는 종종 날카로운 최소값으로 수렴하고 도메인 시프트 하에서 일반화 성능이 떨어지기 때문이다.
  • 이론적으로도 실험적으로도 평탄한 최소값이 도메인 일반화(DG) 상황에서 더 나은 일반화를 이끌어낸다는 것을 입증하기.
  • 구조적 변경 없이도 평탄함과 일반화 성능을 향상시킬 수 있는 단순하면서도 효과적인 방법인 확률적 가중치 평균화를 조밀하게 적용한(SWAD) 방법을 개발하기.
  • 데이터 증강 및 일致성 정규화 방법과 비교함으로써, 성능 향상 요인이 평탄함인지, 내부 도메인 일반화 성능 향상인지 확인하기.

제안 방법

  • 파arameter 공간 이웃 영역 내의 최악의 경험적 리스크를 이용해 도메인 일반화 갭을 상한선으로 제시하는 강건한 리스크 최소화(RRM) 공식을 제안한다.
  • 이론적으로 최적 해 주변의 이웃에서 최악의 리스크를 상한선으로 제시함으로써, 평탄한 최소값이 더 작은 도메인 일반화 갭을 유도한다는 것을 보여준다.
  • 손실 곡면의 평탄한 영역를 더 잘 탐색하기 위해, 모든 학습 반복 단계에서 가중치를 조밀하게 샘플링하는 방식으로 확률적 가중치 평균화(SWA)를 수정한다.
  • 검증 손실을 이용해 과적합을 방지하기 위해 평균화의 최적 시작 및 종료 반복 수를 결정하는 과적합을 고려한 전략을 도입한다.
  • 구조적 변경 없이 기존 DG 방법에 SWAD를 플러그인 방식으로 적용함으로써 일관된 성능 향상을 이룬다.
  • CPU 메모리를 활용해 중간 가중치를 저장함으로써 GPU 메모리 오버헤드를 최소화하면서도 학습 효율성을 유지한다.

실험 결과

연구 질문

  • RQ1비볼록적이고 복잡한 딥 러닝 곡면에서 평탄한 최소값을 찾는 것이 도메인 일반화 갭을 상당히 줄일 수 있는가?
  • RQ2분포 시프트가 i.i.d. 설정보다 더 심각한 상황에서 도메인 시프트 하에서도 평탄한 최소값의 일반화 성능 향상 효과가 유지되는가?
  • RQ3구조적 변경 없이도 복잡한 작업 특화 DG 방법보다 단순한 평탄함 인식 최적화 방법인 SWAD가 뛰어난 성능을 낼 수 있는가?
  • RQ4SWAD의 성능 향상 요인이 내부 도메인 일반화 향상 때문인지, 아니면 더 나은 도메인 외 강건성 때문인지?
  • RQ5도메인 일반화 맥락에서, SWAD는 평탄함 인식 방법(SAM, SWA)과 내부 도메인 일반화 기법(Mixup, CutMix)과 비교해 어떻게 성능를 내는가?

주요 결과

  • SWAD는 다섯 개인 주요 DG 벤치마크(PACS, VLCS, OfficeHome, TerraIncognita, DomainNet)에서 최신 기준 성능(SOTA)을 달성하며, ERM 대비 각각 +2.6pp, +1.6pp, +4.1pp, +3.9pp, +5.6pp 향상된다.
  • 평균적으로 SWAD는 기존 최고의 SOTA 방법 대비 도메인 외 정확도를 1.6%p 향상시키며, ERM 기준으로는 3.6pp 향상된다.
  • 이전 SOTA 방법(SOTA [31])과 SWAD를 조합하면 추가적인 성능 향상이 이루어져 평균 정확도 67.3%를 달성하며, 이는 SWAD 단독 사용 대비 0.4pp 높은 성능이다.
  • 손실 곡면 시각화와 평탄도 지표를 통해, SWAD는 기존의 단순한 SWA보다 항상 더 평탄한 최소값을 찾는 것으로 확인되었다.
  • 데이터 증강 및 일치성 정규화 방법(Mixup, CutMix 등)은 도메인 외 일반화 성능 향상에 기여하지 않았지만, 평탄함 인식 방법(SWA, SAM)은 기여했으며, 이는 평탄도가 핵심 요인임을 확인한다.
  • SWAD는 ERM 대비 런타임 오버헤드가 1.07배에서 1.27배 수준이며, 추가 GPU 메모리 비용이 없어 실세계 적용에 실용적이다.

더 나은 연구,지금 바로 시작하세요

논문 읽기부터 검토까지, 연구 시간을 획기적으로 줄여보세요.

카드 등록 없음 · 무료 플랜 제공

이 리뷰는 AI가 만들고, 인간 에디터가 검토했습니다.