ml.lab
Python sleeps until you run code
07 Recurrent networks

Lesson 5 of 5

Vanishing and exploding gradients

The gradient from k steps back is a product of k factors built from the same recurrent weights, so it shrinks or grows geometrically. Clipping tames the growth; nothing simple fixes the shrinking, and a plain RNN fails a memory test once the delay grows beyond a short range.

About 60 minutes
By the end you can
  • Follow a gradient back through a scalar RNN hop by hop, computing the factor w(1−ht2)w(1 - h_t^2) for each hop.

  • Show that the gradient from kk steps back is a product of kk Jacobians, and bound its size with the largest singular value of the recurrent matrix.

  • Explain why saturated states block the gradient, and why a recurrent weight above 1 does not guarantee growth.

  • Explain why clipping tames exploding gradients but nothing so simple fixes vanishing ones, and describe truncated backpropagation through time.

  • Show with a delayed-recall experiment that a plain RNN fails once the delay grows beyond a short range.

The character RNN from the previous lesson learned to beat the bigram model, and its samples look like Shakespeare line by line. But little carries over from one line to the next. Is that the model's size, the training, or something deeper?

Something deeper. Backpropagation through time gets a loss's gradient to every earlier state, but on the way back the gradient is multiplied at every step by nearly the same factor. That is the repeated multiplication of powers, stability and explosion, and it makes the signal from far back either vanish or explode. This lesson follows one gradient back step by step, bounds its size in general, looks at what can be done about it, and ends with an experiment where a plain RNN fails a simple memory test. That failure is the reason for the LSTM of module 8.

Follow one gradient back

Start with the smallest case, where every factor is a single number. In a scalar RNN, ht=tanh⁡(w ht−1+u xt)h_t = \tanh(w\,h_{t-1} + u\,x_t), the rule from the previous lesson for the gradient arriving from step tt becomes

∂L∂ht−1=w (1−ht2) ∂L∂ht,\frac{\partial L}{\partial h_{t-1}} = w\,(1 - h_t^2)\,\frac{\partial L}{\partial h_t},

when ht−1h_{t-1} affects the loss only through hth_t (it has no output loss of its own). In words: to go back one step, multiply by the recurrent weight ww and by the tanh slope 1−ht21 - h_t^2 at the state being left. Call w(1−ht2)w(1 - h_t^2) the factor of that hop.

The diagram below applies this rule over four steps. The RNN has w=0.8w = 0.8, u=1u = 1 and h0=0h_0 = 0 and reads the pulse x=1,0,0,0x = 1,\allowbreak 0,\allowbreak 0,\allowbreak 0, the fading memory from the first lesson. The loss depends only on the last state, L=h4L = h_4. Boxes hold the states, amber arcs carry the gradient back one hop at a time with each hop's factor written on it, and a chart underneath collects the gradient after 0, 1, 2 and 3 hops on a log scale.

Now the general scalar case. Follow one loss term ℓt\ell_t back kk steps. Each hop contributes its own factor, and the chain rule multiplies them:

∂ℓt∂ht−k=(∏s=t−k+1tw (1−hs2))∂ℓt∂ht.\frac{\partial \ell_t}{\partial h_{t-k}} = \left(\prod_{s=t-k+1}^{t} w\,(1 - h_s^2)\right)\frac{\partial \ell_t}{\partial h_t}.

The product runs over the kk states the gradient leaves on its way back, hth_t down to ht−k+1h_{t-k+1}. In the diagram, t=4t = 4 and k=3k = 3: the product is w(1−h42)⋅w(1−h32)⋅w(1−h22)=0.7200×0.6659×0.5636=0.2702w(1 - h_4^2) \cdot w(1 - h_3^2) \cdot w(1 - h_2^2) = 0.7200 \times 0.6659 \times 0.5636 = 0.2702, the gradient at h1h_1. If the state hovers around some value hh, every factor is about γ=w(1−h2)\gamma = w(1 - h^2), and the product of kk of them is γk\gamma^k. Some numbers:

  • w=0.9w = 0.9 and states near h=0.5h = 0.5: γ=0.9×(1−0.25)=0.9×0.75=0.675\gamma = 0.9 \times (1 - 0.25) = 0.9 \times 0.75 = 0.675. After 10 steps the gradient is 0.67510≈0.01960.675^{10} \approx 0.0196 of its size, after 20 steps 0.67520≈0.0003860.675^{20} \approx 0.000386.
  • w=1.5w = 1.5 and states near h=0.6h = 0.6: γ=1.5×(1−0.36)=1.5×0.64=0.96\gamma = 1.5 \times (1 - 0.36) = 1.5 \times 0.64 = 0.96. Slow decay, but decay: 0.9620≈0.440.96^{20} \approx 0.44 and 0.96100≈0.0170.96^{100} \approx 0.017. A weight above 1 is not enough to keep the gradient alive, because the tanh slope is below 1.
  • w=2w = 2 with the state latched at h=0.9575h = 0.9575 (the first lesson's memory): γ=2×(1−0.9168)=2×0.0832=0.166\gamma = 2 \times (1 - 0.9168) = 2 \times 0.0832 = 0.166. A saturated state is a flat part of the tanh, and the gradient through it dies within a few steps, as step 8 of the diagram showed for w=1.8w = 1.8. The very mechanism that held a memory in the forward pass blocks the learning signal in the backward pass.

The tanh slope is never more than 1, so each factor has size at most ∣w∣|w|. The gradient can only grow if ∣w∣(1−h2)>1|w|(1 - h^2) > 1, which needs ∣w∣>1|w| > 1 and states close enough to 0 that the tanh is steep there.

A beam of amber light enters at the right and travels left through eight identical, evenly spaced panes of tinted glass. Each pane lets through the same fraction of the light that reaches it, so the beam dims by the same ratio at every pane and is only a faint glow at the far left.Each step back is one more pane: the same fraction gets through every time, so the signal from far back fades geometrically.

Walk three hops yourself, through states that are not all alike. One hop grows the gradient and one shrinks it hard; see which and why.

On paperThree hops back

A scalar RNN has w=1.2w = 1.2, and its loss depends only on the last state, L=h4L = h_4, so ∂L/∂h4=1\partial L/\partial h_4 = 1. The forward pass recorded h2=0.1h_2 = 0.1, h3=0.8h_3 = 0.8 and h4=0.5h_4 = 0.5 (the inputs that produced them do not matter here).

  1. Compute the factor w(1−ht2)w(1 - h_t^2) of each hop: from h4h_4 to h3h_3, from h3h_3 to h2h_2, and from h2h_2 to h1h_1.
  2. Compute ∂L/∂h3\partial L/\partial h_3, ∂L/∂h2\partial L/\partial h_2 and ∂L/∂h1\partial L/\partial h_1 by multiplying the factors in turn.
  3. Which hop makes the gradient larger, which shrinks it the most, and why?
  4. Compare ∂L/∂h1\partial L/\partial h_1 with the bound w3w^3.

To check your work, enter three numbers as a column: ∂L/∂h3\partial L/\partial h_3, ∂L/∂h2\partial L/\partial h_2 and ∂L/∂h1\partial L/\partial h_1, to 4 decimals.

Work it on real paper: writing each step is the point. Then check your final answer here and compare your working with the walk-through.

∂L/∂h₃, ∂L/∂h₂, ∂L/∂h₁

One entry per box, top to bottom. 0.25, -2, 3/4 and sqrt(2) all work. Enter moves to the next empty box and checks once all are filled.

When the states stay in one region, one number, γ\gamma, decides everything. Try it for a weight above 1.

Work it outTwenty steps back

A scalar RNN has whh=1.2w_{hh} = 1.2, and during a long stretch of text its state stays near h=0.5h = 0.5. By roughly what factor is a gradient multiplied on its way back 20 steps through that stretch? Give 3 decimal places.

factor

Type a number: 0.25, -2, 3/4 and sqrt(2) all work. Enter checks.

A product of Jacobians

A real RNN has many hidden units, so each factor is a matrix, not a number. The story is the same, with matrices. Between ht−k\mathbf{h}_{t-k} and ht\mathbf{h}_t lie kk recurrent steps, and each step is a function from the previous state to the next. Its Jacobian, the matrix of partial derivatives ∂hs,i/∂hs−1,j\partial h_{s,i}/\partial h_{s-1,j}, is

Js=∂hs∂hs−1=diag(1−hs2) Whh.\mathbf{J}_s = \frac{\partial \mathbf{h}_s}{\partial \mathbf{h}_{s-1}} = \mathrm{diag}(1 - \mathbf{h}_s^2)\,\mathbf{W}_{hh}.

Here diag(v)\mathrm{diag}(\mathbf{v}) is the diagonal matrix with the entries of v\mathbf{v} on its diagonal, so Js\mathbf{J}_s is the recurrent matrix with row ii scaled by the tanh slope 1−hs,i21 - h_{s,i}^2 of unit ii. To see it, differentiate hs,i=tanh⁡(∑j(Whh)ijhs−1,j+⋯)h_{s,i} = \tanh\big(\sum_j (W_{hh})_{ij}h_{s-1,j} + \cdots\big) with respect to hs−1,jh_{s-1,j}: the chain rule gives the tanh slope 1−hs,i21 - h_{s,i}^2 times the inner derivative (Whh)ij(W_{hh})_{ij}. With one unit, Js\mathbf{J}_s is the number w(1−hs2)w(1 - h_s^2) from the previous section.

The chain rule is multiplication of Jacobians, so

∂ht∂ht−k=Jt Jt−1⋯Jt−k+1,∂ℓt∂ht−k=Jt−k+1⊤⋯Jt−1⊤ Jt⊤ ∂ℓt∂ht,\frac{\partial \mathbf{h}_t}{\partial \mathbf{h}_{t-k}} = \mathbf{J}_t\,\mathbf{J}_{t-1}\cdots\mathbf{J}_{t-k+1}, \qquad \frac{\partial \ell_t}{\partial \mathbf{h}_{t-k}} = \mathbf{J}_{t-k+1}^\top\cdots\mathbf{J}_{t-1}^\top\,\mathbf{J}_t^\top\,\frac{\partial \ell_t}{\partial \mathbf{h}_t},

with Js⊤=Whh⊤diag(1−hs2)\mathbf{J}_s^\top = \mathbf{W}_{hh}^\top\mathrm{diag}(1 - \mathbf{h}_s^2) (a diagonal matrix is its own transpose). The second formula is the gradient version. By the chain rule, ∂ℓt/∂ht−k\partial\ell_t/\partial\mathbf{h}_{t-k} is the transpose of the Jacobian product times the gradient at ht\mathbf{h}_t, as in Jacobians and the chain rule, and transposing a product reverses its order: (Jt⋯Jt−k+1)⊤=Jt−k+1⊤⋯Jt⊤(\mathbf{J}_t\cdots\mathbf{J}_{t-k+1})^\top = \mathbf{J}_{t-k+1}^\top\cdots\mathbf{J}_t^\top. Read it from right to left: the gradient starts at ht\mathbf{h}_t and is multiplied by one factor for every step it travels back, each factor being the previous lesson's rule, "through the tanh slopes, then through Whh⊤\mathbf{W}_{hh}^\top". After kk steps it has been multiplied by kk matrices, each built from the same Whh\mathbf{W}_{hh}.

A bound from singular values

With matrices we cannot just multiply numbers, but we can bound how much each factor stretches. In symmetric matrices and the SVD the largest singular value σ1\sigma_1 was the gain of a matrix: ∣Av∣≤σ1(A)∣v∣|\mathbf{A}\mathbf{v}| \le \sigma_1(\mathbf{A})|\mathbf{v}| for every v\mathbf{v}. Each backward factor does two things:

  1. It multiplies the entries of the gradient by the slopes, each between 0 and 1, which cannot make it longer.
  2. It multiplies by Whh⊤\mathbf{W}_{hh}^\top, which has the same singular values as Whh\mathbf{W}_{hh} and so stretches by at most σ1(Whh)\sigma_1(\mathbf{W}_{hh}).

So each step back can stretch the gradient by at most σ1(Whh)\sigma_1(\mathbf{W}_{hh}), and after kk steps

∣∂ℓt∂ht−k∣≤σ1(Whh)k∣∂ℓt∂ht∣.\left|\frac{\partial \ell_t}{\partial \mathbf{h}_{t-k}}\right| \le \sigma_1(\mathbf{W}_{hh})^k\left|\frac{\partial \ell_t}{\partial \mathbf{h}_t}\right| .

If σ1(Whh)<1\sigma_1(\mathbf{W}_{hh}) < 1, the gradient from kk steps back is guaranteed to shrink at least geometrically: vanishing is certain. If σ1>1\sigma_1 > 1 the bound allows growth but does not force it, because the tanh slopes may pull every factor below 1, as they did for w=1.8w = 1.8 in the diagram.

Go slower: The bound, one factor at a time

Let v\mathbf{v} be any gradient vector and d=1−hs2\mathbf{d} = 1 - \mathbf{h}_s^2 the slopes, each in (0,1](0,\allowbreak 1]. The factor is Js⊤v=Whh⊤(d⊙v)\mathbf{J}_s^\top\mathbf{v} = \mathbf{W}_{hh}^\top(\mathbf{d} \odot \mathbf{v}), since multiplying by diag(d)\mathrm{diag}(\mathbf{d}) multiplies entry by entry.

First, the slopes cannot lengthen v\mathbf{v}: ∣d⊙v∣2=∑idi2vi2≤∑ivi2=∣v∣2|\mathbf{d} \odot \mathbf{v}|^2 = \sum_i d_i^2 v_i^2 \le \sum_i v_i^2 = |\mathbf{v}|^2, since every di2≤1d_i^2 \le 1.

Second, Whh⊤\mathbf{W}_{hh}^\top stretches by at most σ1(Whh)\sigma_1(\mathbf{W}_{hh}). If Whh=UΣV⊤\mathbf{W}_{hh} = \mathbf{U}\boldsymbol{\Sigma}\mathbf{V}^\top then Whh⊤=VΣU⊤\mathbf{W}_{hh}^\top = \mathbf{V}\boldsymbol{\Sigma}\mathbf{U}^\top, an SVD with the same singular values. So ∣Whh⊤u∣≤σ1(Whh) ∣u∣|\mathbf{W}_{hh}^\top\mathbf{u}| \le \sigma_1(\mathbf{W}_{hh})\,|\mathbf{u}| for every u\mathbf{u}.

Together, with u=d⊙v\mathbf{u} = \mathbf{d} \odot \mathbf{v}: ∣Js⊤v∣≤σ1(Whh) ∣d⊙v∣≤σ1(Whh) ∣v∣|\mathbf{J}_s^\top\mathbf{v}| \le \sigma_1(\mathbf{W}_{hh})\,|\mathbf{d} \odot \mathbf{v}| \le \sigma_1(\mathbf{W}_{hh})\,|\mathbf{v}|. Apply this kk times, once per factor, starting from v=∂ℓt/∂ht\mathbf{v} = \partial\ell_t/\partial\mathbf{h}_t: each application multiplies the bound by one more σ1(Whh)\sigma_1(\mathbf{W}_{hh}), which gives σ1k\sigma_1^k.

The eigenvalue view

The singular-value bound holds for any states. When the states settle into a steady pattern, we can say more. The factors are then nearly the same matrix J\mathbf{J} every step, and a product of kk copies behaves like Jk\mathbf{J}^k. The eigenvalues of powers, stability and explosion take over: the gradient scales like ρ(J)k\rho(\mathbf{J})^k, where ρ(J)=max⁡i∣λi∣\rho(\mathbf{J}) = \max_i |\lambda_i| is the spectral radius, the size of the largest eigenvalue (J⊤\mathbf{J}^\top has the same eigenvalues as J\mathbf{J}).

Now think about what holding a memory requires. A state that stores something robustly sits at a fixed point that pulls nearby states back toward it, so that small disturbances die out instead of wiping the memory. Pulling nearby states back means the Jacobian shrinks small differences, so ρ(J)<1\rho(\mathbf{J}) < 1, and the gradient, which travels back through the same Jacobians, shrinks too. Storing information robustly and passing a learning signal back through it pull in opposite directions. The scalar latch above is the one-unit version: its factor 0.1660.166 is also how much a small disturbance of the stored state shrinks at each step, which is what makes the memory robust and the gradient vanish.

Watching it happen

The instrument below runs the matrix version on a random RNN with 8 tanh units. Its recurrent matrix is Whh=ρ Q\mathbf{W}_{hh} = \rho\,\mathbf{Q}, where Q\mathbf{Q} is a fixed random orthogonal matrix (every singular value 1), so ρ\rho is both the spectral radius and the largest singular value of Whh\mathbf{W}_{hh}. The chart plots ∥∂hT/∂hT−k∥\|\partial\mathbf{h}_T/\partial\mathbf{h}_{T-k}\| against the number of steps back kk, from 0 to 80: the spectral norm of the whole product of kk Jacobians, written with double bars as in symmetric matrices and the SVD, the most it can stretch any gradient. The vertical axis is a log scale, so geometric decay or growth is a straight line. A toggle includes the tanh slopes or leaves them out.

The same experiment in code, with a 64-unit RNN driven by random inputs. G is a random matrix with spectral radius close to 1 (and largest singular value close to 2), and W_hh is G times a scale. The code walks a gradient of length 1 back from the last state with the previous lesson's rule and records its length after every step.

⌘+Enter runs · edit freelyPython sleeps until you run code

At scales 0.5 and 1 the gradient from 80 steps back is smaller than 10−1310^{-13} of where it started. At scale 4 it is larger than 101210^{12}. Only in a narrow band around scale 2 does it stay within a few powers of ten, and nothing in training keeps Whh\mathbf{W}_{hh} inside that band. Notice also which way the singular-value bound works. At scale 1 the largest singular value is about 2, so the bound allows growth by up to 280≈10242^{80} \approx 10^{24}, yet the gradient vanished: the tanh slopes and the directions of the matrices did the rest. A bound below 1 guarantees vanishing; a bound above 1 guarantees nothing.

To recap: going back kk steps multiplies the gradient by kk factors built from the same Whh\mathbf{W}_{hh} and the tanh slopes; their size is at most σ1(Whh)k\sigma_1(\mathbf{W}_{hh})^k; saturated states make the factors smaller still; and in practice the result is geometric decay or growth.

What vanishing does to learning

Why does a tiny gradient from far back matter, if the gradient from nearby steps is fine? Because every weight update adds them together. Every entry of ∂L/∂Whh=∑tδtht−1⊤\partial L/\partial\mathbf{W}_{hh} = \sum_t \boldsymbol{\delta}_t\mathbf{h}_{t-1}^\top is a sum of contributions, and each contribution links a loss at some step to a state some number of steps kk earlier, scaled by roughly γk\gamma^k.

With γ=0.7\gamma = 0.7, a link one step long counts 0.70.7, a link five steps long counts 0.75≈0.170.7^5 \approx 0.17, and a link 20 steps long counts 0.720≈0.00080.7^{20} \approx 0.0008. The long-range contributions are in the sum, but they are drowned out by the short-range ones. Gradient descent follows the sum, so it learns the short-range patterns and barely sees the long-range ones.

What can be done

The two problems are not symmetric, and neither are their remedies.

Exploding gradients: clip them. This is the clipping of the previous lesson. When the product of factors blows up, clipping by norm caps the length of the whole gradient and keeps its direction, so one bad batch cannot throw the weights far away. It treats the symptom, a step that is too long, and that is enough in practice.

Vanishing gradients: no simple fix. Clipping cannot help, and neither can scaling the gradient up, because the problem is relative: the long-range contributions are tiny compared with the short-range ones in the same sum, and scaling multiplies both. Optimizers such as Adam, which divide each weight's step by the typical size of its gradient, help when a gradient is small but consistent; they cannot pull a faint long-range signal out from under larger short-range contributions to the same weights. Careful initialization (a recurrent matrix close to the identity, or an orthogonal one, whose singular values are all 1) helps at the start of training, but nothing keeps the matrix there. The fix that worked was a change of architecture: give the network a memory path along which the per-step factor is a number the network controls and can hold near 1. That is the cell state of the LSTM, the subject of the next module.

Truncated backpropagation through time. A different, practical limit also cuts long-range learning. A long text cannot be unrolled all at once: the forward pass stores every state, so memory grows with every step. The usual compromise cuts the text into consecutive chunks of TT steps and treats the two directions differently:

  • The forward pass carries the state across chunk boundaries: the last state of one chunk becomes h0\mathbf{h}_0 of the next, so the model can still carry information forward indefinitely.
  • The backward pass stops at the boundary: the carried-in h0\mathbf{h}_0 is treated as a constant, so no gradient crosses into the previous chunk.

The model can use information from further back than TT steps, but it gets no training signal telling it to keep such information. The training in the previous lesson was simpler still: every window was cut from a random place and started from a zero state, so the model never saw more than 32 characters of context while it trained.

Check that you can tell apart what truncation keeps and what it cuts.

Quick checkTruncated backpropagation through time

You train on a long text with truncated BPTT: chunks of 50 characters, the last hidden state of each chunk carried into the next chunk as its starting state, and gradients stopped at chunk boundaries. Which statement is true?

Choose one answer, then check.

Where the plain RNN breaks: delayed recall

Everything so far says the long-range signal is weak. A clean experiment shows how weak, with a task that needs memory and nothing else. The delayed-recall task uses 9 tokens: four symbols (ids 0 to 3), four noise tokens (ids 4 to 7) and a cue (id 8). Each sequence is

s, n1, n2, …, nd, cue, ss,\ n_1,\ n_2,\ \ldots,\ n_d,\ \text{cue},\ s

where ss is a random symbol, the nin_i are dd random noise tokens, and dd is the delay. For example, with d=3d = 3 one sequence might be 2, 5, 7, 4, 8, 2: the symbol 2, three noise tokens, the cue, and the symbol again.

The model is trained exactly like the language model: predict the next token at every position, with teacher forcing and the average cross-entropy. Most targets are noise tokens, which nothing can predict better than a one-in-four guess. The one that matters is the last: the token after the cue is ss, which was read d+1d + 1 steps earlier. Guessing gives 25% accuracy at that position; remembering gives 100%.

This is the previous sections in miniature. The gradient that could teach the network to store ss comes from a single position and has to travel back d+1d + 1 steps, while the noise positions contribute larger, short-range gradients to the same weights.

Code itDelayed recall: how far back can it remember

Build the delayed-recall task and measure how far back a plain RNN can remember.

  1. make_recall_batch(batch_size, delay, rng): sequences s,n1,…,nd,cue,ss,\allowbreak n_1,\allowbreak \ldots,\allowbreak n_d,\allowbreak \text{cue},\allowbreak s, turned into next-token pairs.
  2. recall_accuracy(params, delay, rng, n): the fraction of fresh sequences whose last prediction is ss.
  3. train_recall(delay, ...): train a fresh RNN with the language-model loss at every position.
  4. delay_sweep(delays): one trained RNN per delay, and its accuracy.

Provided: init_rnn, rnn_forward, rnn_loss_and_grads and clip_grads from mlref.rnn, Adam from mlref.nn, and the constants K = 4 symbols, N_NOISE = 4 noise tokens, CUE = 8 and V = 9. The tests run the sweep over delays 2, 4, 8, 16, 24 and 32, which takes about 10 seconds.

This one trains a model in your browser, so a run can take a minute or more. The charts update as it goes, and Stop ends it early.

⌘+Enter runsPython sleeps until you run code
Write your code where the starter says raise NotImplementedError, then press Run tests. Each check says what it expects.

In our runs the RNN solves delays 2, 4 and 8 completely. At 16 the outcome depends on the run: with some random seeds it learns the task, with others it gets partway or stays at chance. At 24 and 32 it stays at chance, 25%: after 800 training steps its answers show no trace of ss. Training longer moves the boundary out only slowly, because each extra step of delay multiplies the useful signal by another factor below 1. Module 8 runs a version of this experiment, set up as a classifier that answers once at the end, with an RNN and an LSTM side by side, and the LSTM keeps learning at delays where the RNN usually fails.

To finish the module, explain the whole chain of reasoning in your own words, from the product of factors to the fix.

In your own wordsWhy plain RNNs forget

Explain to a colleague why a plain RNN trained with gradient descent struggles to learn dependencies more than a few dozen steps back. Cover where the product of Jacobians comes from, what bounds its size, why gradient clipping does not help, and what kind of fix does.

A few sentences first: 0 of 60 characters.

Saved in this browser as you type.

Next: the module checkpoint