๐Ÿ”— Deep Learning ยท Lecture 38 of 38

Loss Landscapes, Saddle Points and Flat Minima

What does the surface that SGD descends actually look like? We study critical points in high dimensions, visualise loss landscapes, discuss sharp versus flat minima, mode connectivity and why architecture shapes trainability.

Training a neural network is a descent across a loss surface with millions or billions of dimensions. We cannot see it directly, yet its geometry determines whether training succeeds, how fast it converges and how well the result generalises. Over the past decade, researchers have developed surprising insights into this geometry โ€” many of which overturn low-dimensional intuitions.

Critical points in high dimensions#

A critical point has zero gradient. Its type is determined by the eigenvalues of the Hessian:

  • all positive โ†’ local minimum;
  • all negative โ†’ local maximum;
  • mixed signs โ†’ saddle point.

In low dimensions we picture landscapes full of local minima. In high dimensions, a random critical point needs every one of its millions of eigenvalues to be positive to be a minimum โ€” an exponentially unlikely event unless the loss is already low. Dauphin et al. (2014) argued, drawing on random-matrix and statistical-physics results, that:

  • saddle points vastly outnumber local minima in high dimensions;
  • local minima tend to have loss close to the global minimum;
  • high-loss critical points are mostly saddles, which slow training (plateaus) but can be escaped.

Gradient noise in SGD and momentum help escape saddles; theory shows perturbed gradient descent escapes strict saddles efficiently.

Visualising loss landscapes#

We cannot plot millions of dimensions, but we can plot slices. Li et al. (2018) evaluated the loss along two random directions $\boldsymbol{\delta}, \boldsymbol{\eta}$ around trained parameters $\boldsymbol{\theta}^*$:

$$ f(\alpha, \beta) = L(\boldsymbol{\theta}^* + \alpha\boldsymbol{\delta} + \beta\boldsymbol{\eta}) $$

with filter normalisation (scaling each random direction's filters to match the trained filters' norms, removing scale artefacts). Their striking finding: deep networks without skip connections have chaotic, highly non-convex landscapes, while ResNets and wider networks have smooth, nearly convex-looking basins โ€” a visual explanation of why residual connections make training easier.

python
import torch
import numpy as np

def loss_slice(model, loss_fn, data, steps=21, span=1.0):
    """Evaluate the loss along one filter-normalised random direction."""
    theta = [p.detach().clone() for p in model.parameters()]
    direction = []
    for p in theta:
        d = torch.randn_like(p)
        if p.dim() > 1:                                   # filter/row-wise normalisation
            d = d * (p.norm(dim=tuple(range(1, p.dim())), keepdim=True) /
                     (d.norm(dim=tuple(range(1, p.dim())), keepdim=True) + 1e-10))
        else:
            d = torch.zeros_like(p)                        # common choice: ignore biases/BN params
        direction.append(d)
    alphas, losses = np.linspace(-span, span, steps), []
    with torch.no_grad():
        for a in alphas:
            for p, t, d in zip(model.parameters(), theta, direction):
                p.copy_(t + a * d)
            losses.append(sum(loss_fn(model(x), y).item() for x, y in data) / len(data))
        for p, t in zip(model.parameters(), theta):
            p.copy_(t)                                     # restore trained weights
    return alphas, losses

Sharp versus flat minima#

Hochreiter and Schmidhuber (1997) proposed that flat minima โ€” wide regions where the loss stays low โ€” generalise better than sharp minima, because a flat solution is robust to the differences between the training and test loss surfaces and needs less precision to describe (a minimum-description-length argument). Keskar et al. (2017) reported that large-batch training tended to converge to sharper minima and generalise worse than small-batch training.

Caveats: Dinh et al. (2017) showed that sharpness is not invariant to reparameterisation โ€” for ReLU networks, rescaling weights between layers can make a minimum arbitrarily sharp without changing the function. So flatness must be measured carefully (e.g. with scale-invariant definitions). Nevertheless, methods that explicitly seek flat regions help in practice:

  • Sharpness-Aware Minimisation (SAM) minimises the worst-case loss in a small neighbourhood:
$$ \min_{\boldsymbol{\theta}}\;\max_{\|\boldsymbol{\epsilon}\| \le \rho}L(\boldsymbol{\theta} + \boldsymbol{\epsilon}) $$

approximated with one extra gradient step, often improving generalisation.

  • Stochastic Weight Averaging (SWA) averages weights along the trajectory, landing in the centre of a wide basin.

Mode connectivity#

Train two networks from different random seeds and they converge to different minima. Are these isolated valleys? Garipov et al. (2018) and Draxler et al. (2018) found they can be joined by simple curves (e.g. a quadratic Bรฉzier curve) along which the loss stays low: the minima are connected by low-loss paths. Further work found that, after accounting for permutation symmetries of hidden units (neurons can be reordered without changing the function), many solutions are connected even by straight lines (linear mode connectivity), especially in wide networks. This geometry underpins model merging and model soups โ€” averaging the weights of fine-tuned models to combine their strengths.

The edge of stability#

Cohen et al. (2021) observed that with full-batch gradient descent, the largest Hessian eigenvalue (sharpness) rises during training until it reaches about $2/\eta$ โ€” the classical stability limit โ€” and then hovers there, while the loss keeps decreasing non-monotonically. Training operates at the edge of stability, contradicting the assumption that step sizes stay safely below $2/\lambda_{\max}$. The learning rate therefore implicitly controls the sharpness of the solutions found.

Practical takeaways#

JA
Written by

Janin A Apurba

B.Sc. in CSE, AUST ยท Advanced ICT Officer, CNRS-UNHCR. Teaching AI, ML and Deep Learning to the next generation of engineers and researchers.

Keep learning

Related lectures

๐Ÿ”— Deep Learning

Optimisers II: AdaGrad, RMSProp, Adam and AdamW

Adaptive optimisers give each parameter its own learning rate. We derive AdaGrad, RMSProp and Adam including bias correction, explain why AdamW decouples weight decay, and survey newer optimisers.

Intermediateโฑ 6 min#105
๐Ÿ”— Deep Learning

Optimisers I: SGD, Momentum and Nesterov Acceleration

Plain SGD zig-zags through ravines and crawls across plateaus. Momentum accumulates velocity to fix both. We derive heavy-ball and Nesterov momentum, analyse their effect on ill-conditioned problems, and give tuning advice.

Intermediateโฑ 5 min#104
๐Ÿ”— Deep Learning

Double Descent and the Generalisation Mystery of Deep Learning

Over-parameterised networks can fit random labels yet generalise on real data, and test error can fall again beyond the interpolation threshold. We explore double descent, benign overfitting and implicit regularisation.

Advancedโฑ 6 min#133