Skip to main content
QUICK REVIEW

[논문 리뷰] A Riemannian Block Coordinate Descent Method for Computing the Projection Robust Wasserstein Distance

Minhui Huang, Shiqian Ma|arXiv (Cornell University)|2020. 12. 09.
Stochastic Gradient Optimization Techniques참고 문헌 29인용 수 8
한 줄 요약

이 논문은 리만 다성분 좌표 강하(Riemannian Block Coordinate Descent, RBCD) 방법을 제안하여, 스티펠 만류상에서의 비볼록 최대-최소 문제로 재구성된 투영 강건 워샤흐스타인(PRJ) 거리의 효율적 계산을 위해 제안한다. 이 방법은 기존의 RGAS 알고리즘의 $O(\epsilon^{-12})$ 복잡도에 비해 훨씬 향상된 $O(\epsilon^{-3})$ 반복 복잡도를 달성하며, MNIST와 같은 대규모 데이터셋에서 뛰어난 효율성을 보여준다.

ABSTRACT

The Wasserstein distance has become increasingly important in machine learning and deep learning. Despite its popularity, the Wasserstein distance is hard to approximate because of the curse of dimensionality. A recently proposed approach to alleviate the curse of dimensionality is to project the sampled data from the high dimensional probability distribution onto a lower-dimensional subspace, and then compute the Wasserstein distance between the projected data. However, this approach requires to solve a max-min problem over the Stiefel manifold, which is very challenging in practice. The only existing work that solves this problem directly is the RGAS (Riemannian Gradient Ascent with Sinkhorn Iteration) algorithm, which requires to solve an entropy-regularized optimal transport problem in each iteration, and thus can be costly for large-scale problems. In this paper, we propose a Riemannian block coordinate descent (RBCD) method to solve this problem, which is based on a novel reformulation of the regularized max-min problem over the Stiefel manifold. We show that the complexity of arithmetic operations for RBCD to obtain an $ε$-stationary point is $O(ε^{-3})$. This significantly improves the corresponding complexity of RGAS, which is $O(ε^{-12})$. Moreover, our RBCD has very low per-iteration complexity, and hence is suitable for large-scale problems. Numerical results on both synthetic and real datasets demonstrate that our method is more efficient than existing methods, especially when the number of sampled data is very large.

연구 동기 및 목표

  • 스티펠 만류상에서의 비볼록 최대-최소 공식화로 인해 계산에 어려움을 겪는 투영 강건 워샤흐스타인(PRJ) 거리 계산의 계산적 과제를 해결한다.
  • 각 반복에서 엔트로피 정규화된 옵티멀 트랜스포트 문제를 해결해야 하는 기존 방법들, 예를 들어 RGAS와 같은 방법들의 높은 반복당 비용을 극복한다.
  • 산술 복잡도와 반복당 비용을 줄여 대규모 최적화 운반 문제를 위한 더 효율적인 알고리즘을 개발한다.
  • 리만 다성분에서의 제안된 RBCD 방법에 대한 이론적 수렴 보장을 제공한다.
  • 합성 및 실제 데이터셋에서 런타임과 확장성 측면에서 RBCD가 RGAS 및 기타 기준 대비 경험적으로 뛰어난 성능을 보임을 입증한다.

제안 방법

  • PRW 계산을 위한 정규화된 최대-최소 문제를, 스티펠 만류상에서 블록 좌표 강하에 적합한 새로운 형태로 재구성한다.
  • 스티펠 만류상 $\mathrm{St}(d,k)$ 상에서 운반 계획 $\pi$ 와 투영 행렬 $U$ 에 대해 번갈아가며 최적화하기 위해 리만 블록 좌표 강하(RBCD)를 적용한다.
  • 업데이트 과정에서 $U^\top U = I_k$ 의 만료 제약 조건을 다루기 위해 리만 최적화 기법을 사용한다.
  • 수렴 속도 향상을 위해 스텝 크기 선택을 위한 선형 탐색을 통한 적응형 변형인 RABCD를 도입한다.
  • 목적 함수 $f(\pi, U) = \sum_{i,j} \pi_{ij} \|U^\top x_i - U^\top y_j\|^2$ 의 구조를 활용하여 하위 문제의 효율적 해법을 가능하게 한다.
  • 하위 문제를 해결하기 위해 리만 공액 기울기 또는 신뢰 영역 방법을 통합하여 낮은 반복당 비용을 확보한다.

실험 결과

연구 질문

  • RQ1리만 블록 좌표 강하 방법이 PRW 거리 계산에 있어 기존 방법들, 예를 들어 RGAS와 비교해 더 낮은 반복 복잡도를 달성할 수 있는가?
  • RQ2제안된 RBCD 방법은 대규모 데이터셋에서 계산 비용을 크게 줄이면서도 정확도를 유지하는가?
  • RQ3다양한 데이터 차원과 표본 크기에서 RBCD 알고리즘이 RGAS 대비 수렴 속도와 확장성 측면에서 어떻게 성능을 내는가?
  • RQ4비볼록 최대-최소 문제에 대해 스티펠 만류상에서 $\epsilon$-정류점에 도달하기 위한 RBCD의 이론적 반복 복잡도는 무엇인가?
  • RQ5제안된 방법은 MNIST를 이용한 숫자 분류와 같은 실제 기계 학습 작업에 효과적으로 적용될 수 있는가?

주요 결과

  • RBCD 알고리즘은 $\epsilon$-정류점에 도달하기 위해 $O(\epsilon^{-3})$ 반복 복잡도를 달성하며, 이는 기존 RGAS 알고리즘의 $O(\epsilon^{-12})$ 복잡도에 비해 상당한 향상이다.
  • 합성 및 실제 데이터셋, 특히 MNIST에서의 수치 실험 결과, RBCD는 표본 수가 증가할수록 RGAS보다 빠르게 작동하는 것으로 나타났다.
  • MNIST 데이터셋에서, RBCD는 모든 숫자 쌍 비교에서 RGAS 대비 최대 30% 런타임을 단축시켰으며, 일관된 PRW 거리 값이 유지되었다.
  • 모든 숫자 쌍에서 RBCD와 RGAS가 계산한 PRW 거리는 매우 유사했으며, 제안된 방법의 신뢰성과 일관성을 확인했다.
  • 적응형 변형인 RABCD는 수렴 속도를 추가로 향상시켰지만, RBCD와 RABCD 모두 다양한 파rameter 설정에서 견고한 성능을 보였다.
  • 이 방법은 파rameter 조정에 대해 뛰어난 내성성을 보였지만, 최적의 성능을 내기 위해서는 스텝 크기 및 정규화 파rameter의 신중한 선택이 필요하다.

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

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

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

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