How to escape sharp minima with random perturbations

Fuente: arXiv
Saved in:
Bibliographic Details
Main Authors: Ahn, Kwangjun, Jadbabaie, Ali, Sra, Suvrit
Format: Preprint
Published: 2023
Subjects:
Online Access:
Tags: Add Tag
No Tags, Be the first to tag this record!
_version_ 1866914810669563904
author Ahn, Kwangjun
Jadbabaie, Ali
Sra, Suvrit
author_facet Ahn, Kwangjun
Jadbabaie, Ali
Sra, Suvrit
contents Modern machine learning applications have witnessed the remarkable success of optimization algorithms that are designed to find flat minima. Motivated by this design choice, we undertake a formal study that (i) formulates the notion of flat minima, and (ii) studies the complexity of finding them. Specifically, we adopt the trace of the Hessian of the cost function as a measure of flatness, and use it to formally define the notion of approximate flat minima. Under this notion, we then analyze algorithms that find approximate flat minima efficiently. For general cost functions, we discuss a gradient-based algorithm that finds an approximate flat local minimum efficiently. The main component of the algorithm is to use gradients computed from randomly perturbed iterates to estimate a direction that leads to flatter minima. For the setting where the cost function is an empirical risk over training data, we present a faster algorithm that is inspired by a recently proposed practical algorithm called sharpness-aware minimization, supporting its success in practice.
format Preprint
id arxiv_https___arxiv_org_abs_2305_15659
institution arXiv
publishDate 2023
record_format arxiv
spellingShingle How to escape sharp minima with random perturbations
Ahn, Kwangjun
Jadbabaie, Ali
Sra, Suvrit
Machine Learning
Artificial Intelligence
Optimization and Control
Modern machine learning applications have witnessed the remarkable success of optimization algorithms that are designed to find flat minima. Motivated by this design choice, we undertake a formal study that (i) formulates the notion of flat minima, and (ii) studies the complexity of finding them. Specifically, we adopt the trace of the Hessian of the cost function as a measure of flatness, and use it to formally define the notion of approximate flat minima. Under this notion, we then analyze algorithms that find approximate flat minima efficiently. For general cost functions, we discuss a gradient-based algorithm that finds an approximate flat local minimum efficiently. The main component of the algorithm is to use gradients computed from randomly perturbed iterates to estimate a direction that leads to flatter minima. For the setting where the cost function is an empirical risk over training data, we present a faster algorithm that is inspired by a recently proposed practical algorithm called sharpness-aware minimization, supporting its success in practice.
title How to escape sharp minima with random perturbations
topic Machine Learning
Artificial Intelligence
Optimization and Control
url https://arxiv.org/abs/2305.15659