Plain gradient descent has no memory. Every step depends only on the gradient at the current point. On a flat region it takes tiny steps, and in a narrow valley it bounces between walls. Momentum adds memory.
The idea
Picture a heavy ball rolling down the surface. It does not stop and restart at each point. It carries velocity, which builds up when the slope keeps pointing the same way and is damped when the slope keeps changing sign.
We keep a velocity vector that accumulates past gradients:
The number , typically 0.9, says how much of the old velocity survives. With this is ordinary gradient descent.
Unrolling the recursion shows what is:
It is an exponentially weighted sum of past gradients, recent ones counting most. Two effects follow:
- If the gradient keeps the same sign, the terms add up. In the long run the step is about , which is for .
- If the gradient alternates in sign, as it does across a valley, the terms largely cancel, which damps the zigzag.
Where it helps: the flat start
The sigmoid toy problem starts on a plateau. With the same learning rate, the number of steps until the loss stays below is:
| Rate | Plain GD | Momentum () |
|---|---|---|
| 0.1 | 1815 | 177 |
| 0.3 | 606 | 124 |
| 1.0 | 183 | 148 |
With a small rate, momentum is about ten times faster: the velocity builds up across the flat stretch. As plain gradient descent gets a good rate the advantage shrinks, because it is no longer crawling.
Where it can hurt: overshoot
A ball with velocity does not stop at the bottom. It rolls through and climbs the other side. On the narrow valley with :
| Step | Momentum | Plain GD |
|---|---|---|
| 0 | (1, 1) | (1, 1) |
| 2 | (−0.20, 0.86) | (0.25, 0.90) |
| 4 | (−0.84, 0.58) | (0.06, 0.81) |
| 6 | (0.03, 0.25) | (0.02, 0.74) |
Momentum makes much more progress along the flat direction: is 0.25 after six steps against 0.74. The price is that swings past zero and back, since the accumulated velocity carries it across. Overall on this small valley, momentum needs 102 steps to settle below against 84 for plain gradient descent at the same rate. Momentum is not a guaranteed improvement. It helps most when the plain method is crawling, and it needs a rate and that do not overshoot too much.
In code
def momentum_step(theta, v, grad, eta=0.05, gamma=0.9):
v = gamma * v + eta * grad(theta)
return theta - v, v
Start with and call it repeatedly. The velocity is the only extra state.
Compare gradient descent and momentum on the plateau and in the narrow valley.