Skip to main content

Stochastic gradient descent

Examples

Stochastic gradient descent

In every step of stochastic gradient descent (SGD) the gradient is computed on a randomly selected subset of the data, a batch, instead of on all of it. As long as the batch is somewhat representative of the full data set, its gradient points more or less in the same direction as the full gradient, and it is much cheaper to compute.

Here the 40 points of the linear regression on gradient descent are split into four batches of ten. SGD uses them one after the other, in the same order every time. One pass through all the batches is an epoch, here four steps.

Top left, the data coloured by batch, the current batch highlighted, and the fit at step \(t\). Top right, the level lines of the loss on all data (grey) and on each batch (coloured, the current one highlighted), with the path of gradient descent (black) and of SGD (red). Bottom, the training loss against the number of data points the gradients have used so far: a gradient descent step uses 40, an SGD step 10. For SGD, the loss on each batch, and the epoch loss at the end of every epoch.

The gradient on a batch is perpendicular to the level lines of that batch’s loss, not to those of the full loss. With a new batch in every step, the path jitters around the path of gradient descent.

For the same learning rate, SGD decreases the training loss faster for the same number of data points used. In practice the loss on all data is not computed in every step of SGD. Instead, training reports the epoch loss, the mean of the losses on the batches of an epoch, which costs nothing extra.

Improved versions of gradient descent

There are many tricks to improve on plain (stochastic) gradient descent. A popular one is momentum, which keeps a running average of past gradients and steps along it (why momentum really works). We do not discuss these ideas further here, but Adam and its variant AdamW are particularly popular and successful improvements. They usually need little or no tuning of the learning rate.

Adam keeps, for every parameter, a running average of its gradients and of their squares, and divides the step by the square root of the latter. Each parameter then moves by a step of similar size, however steep the loss is in its direction. This helps when the loss is much steeper in some directions than in others. Here the linear regression of gradient descent has its input in other units, with a standard deviation of 10 instead of 1, so the loss is about a hundred times steeper in the slope \(\beta_1\) than in the intercept \(\beta_0\).

In torch an optimizer object carries out the update, so the loop no longer subtracts the gradient itself.

import numpy as np
import torch

def advanced_gradient_descent(loss, x, optimizer, T):
    learning_curve = []
    for t in range(T):
        optimizer.zero_grad()
        l = loss(x)
        learning_curve.append(l.item())
        l.backward()
        optimizer.step()
    return x, learning_curve


rng = np.random.default_rng(6)
X = torch.tensor(10 * rng.standard_normal(40))
y = 0.1 - 0.04 * X + 0.1 * torch.tensor(rng.standard_normal(40))

beta = torch.tensor([2.5, 0.8], dtype=torch.float64, requires_grad=True)
optimizer = torch.optim.Adam([beta], lr=0.2)
beta, curve = advanced_gradient_descent(
    lambda b: torch.mean((y - b[0] - b[1] * X) ** 2), beta, optimizer, 200)
print("Adam after 200 steps:", beta.detach().numpy().round(3))
1
Resets the gradients, as x.grad.zero_() did.
2
Updates the parameters it was given, here with the Adam rule.
3
The default learning rate of Adam in torch is \(10^{-3}\); here it is 0.2.
Adam after 200 steps: [ 0.096 -0.039]
Gradient descent and Adam from the same start. Top left, the data and both fits at step \(t\). Bottom left, how far each loss is above its minimum. Right, the level lines of the loss with both paths. Gradient descent uses 0.9 times its largest stable learning rate, Adam the learning rate 0.2. The switch standardizes the input.

With the input as measured, gradient descent needs a learning rate below 0.0099, or it diverges along the steep direction. With such a small learning rate it zigzags across the narrow valley and then crawls along it: after 200 steps its loss is still \(5 \cdot 10^{-3}\) above the minimum, against \(6 \cdot 10^{-9}\) for Adam. With the input standardized, the loss is equally steep in both directions, and gradient descent reaches the minimum faster than Adam.

Early stopping

In early stopping we start (stochastic) gradient descent with small parameter values, keep track of the training and the validation loss throughout, and stop when the validation loss is smallest. The effect is similar to regularization: the parameters found at early stopping usually have a smaller norm than the ones with the lowest training loss.

A polynomial of degree 12, fitted by Adam to 10 points of \(f(x) = 0.3\sin(10x) + 0.7x\) with noise of standard deviation 0.1, from parameters close to zero. Left, the fit at step \(t\). Right, the training loss and the loss on 50 validation points, with their minimum marked.

The validation loss is smallest at step 6533, 0.0142, and rises to 0.0199 by step 30000 while the training loss keeps falling. The norm of the parameters grows from 25.6 at the minimum of the validation loss to 40.8 at the end.