Course Content
Natural Language Processing Basics
4 sections · 10 lessons
LSTMs & GRUs — Solving the Vanishing Gradient Problem
Train a plain recurrent network on product reviews and then feed it this one:
"I ordered this for my daughter's birthday. The packaging arrived dented but the courier was apologetic, and honestly the design is lovely — she was thrilled when she opened it, and for about a week it was the best thing in the house. Then the hinge snapped clean off and the company will not answer my emails."The model says positive, confidently. Look at where its evidence is. Words like lovely, thrilled, best sit in the middle. The verdict — snapped, will not answer — is at the end. And a review whose ending should dominate has been classified on its middle.
Now flip the review so the complaint comes first and the praise last. The model flips too. It is not reading the review. It is reading the last dozen words.
Why the memory horizon exists
The reason is arithmetic and it is worth being precise about, because everything in this lesson is a response to it.
In a plain RNN, the hidden state is completely rewritten at every step:
To learn that a word 60 steps back mattered, the gradient must travel back through 60 of these updates. Each hop multiplies by another Jacobian:
Two things guarantee that this shrinks. The tanh derivative 1−ht2 is at most 1 and drops towards 0 as the state saturates — a state value of 0.9 gives a derivative of 0.19. And well-conditioned weight matrices typically have norm below 1, or training diverges.
So you are multiplying numbers smaller than 1, sixty times over. At a per-step factor of 0.7, the gradient reaching step t−60 is 0.760≈5×10−10 of the gradient at step t.
The plain RNN does not fail to remember. It fails to learn that remembering was worth doing, because the signal telling it so has been multiplied away before it arrives.
Notice the shape of the problem: the damage comes from repeated multiplication. If information could travel along a path that adds instead of multiplies, the gradient would survive. That is the entire idea behind the LSTM.
The LSTM: a separate memory with gates on it
The Long Short-Term Memory cell keeps two vectors instead of one:
- Cell state ct — the long-term memory. It is modified by addition, not by being overwritten.
- Hidden state ht — the working output, a filtered view of ct, and what the next layer sees.
Three gates control what happens to the memory. A gate is a vector of numbers between 0 and 1, produced by a sigmoid, that multiplies another vector element by element. A gate value near 0 blocks; near 1 lets through. Crucially, the gates are computed from the current input and previous state, so the network learns when to open and close them.
The equations, one at a time
Every gate has the same form — look at the current word and the running summary, then squash to (0,1):
"How much of each element of the existing memory should I keep?" On reading a full stop and a new subject, the network can learn to push this towards 0 and clear stale state.
"How much of the newly proposed information should I write in?"
The proposed new content, in (−1,1). Note this uses tanh, not sigmoid — it is content, not a gate.
This is the line that matters. Old memory, scaled by the forget gate, plus new content, scaled by the input gate. ⊙ is element-wise multiplication.
The memory is not the output. The cell can hold something in ct for fifty steps while keeping ot near zero, revealing it only when it becomes relevant.
Why this rescues the gradient
Differentiate the memory update with respect to the previous cell state:
Compare this with the plain RNN's Whh⊤diag(1−ht2). There is no weight matrix. There is no tanh derivative. The gradient along the cell state is multiplied by the forget gate and nothing else.
If the network learns ft≈1 for the dimensions holding something important, the gradient passes through essentially unchanged, however many steps it travels. Put numbers on it:
| Steps back | Plain RNN, factor 0.7 | LSTM, ft=0.99 | LSTM, ft=0.9 |
|---|---|---|---|
| 10 | 0.028 | 0.904 | 0.349 |
| 50 | 1.8×10−8 | 0.605 | 0.005 |
| 100 | 3.2×10−16 | 0.366 | 2.7×10−5 |
| 200 | 1.0×10−31 | 0.134 | 7.1×10−10 |
Read the third column. After 200 steps the gradient retains 13% of its magnitude — perfectly trainable. Read the fourth column and notice the honest caveat: an LSTM with a forget gate at 0.9 is barely better than a plain RNN. The LSTM does not eliminate vanishing gradients. It gives the network a mechanism to avoid them, which it must learn to use.
The gradient path is often called the constant error carousel: a route from ct back to ct−1 that involves no matrix multiplication at all.
One timestep, worked
Take a single dimension so the numbers stay readable. Suppose the cell is tracking "the overall verdict of this review", currently positive.
state before: c = 0.80 (a positive verdict is held in memory)Reading a neutral word ("the"): f = 0.97 keep almost everything i = 0.05 write almost nothing c~ = 0.10 c_new = 0.97*0.80 + 0.05*0.10 = 0.776 + 0.005 = 0.781 o = 0.20 expose little h = 0.20 * tanh(0.781) = 0.20 * 0.653 = 0.131Reading "snapped": f = 0.15 the verdict is being overturned - clear the old value i = 0.92 write the new evidence hard c~ = -0.85 c_new = 0.15*0.781 + 0.92*(-0.85) = 0.117 - 0.782 = -0.665 o = 0.88 this matters now - expose it h = 0.88 * tanh(-0.665) = 0.88 * (-0.582) = -0.512On a neutral word the memory barely moves. On decisive evidence the forget gate closes, the input gate opens, and the verdict flips in a single step. Nothing hand-coded that behaviour; the gate weights were learned from labelled examples.
The bias trick that actually matters
At initialisation, gate biases are usually zero, so σ(0)=0.5 and the forget gate starts at 0.5. From the table above, a factor of 0.5 vanishes fast — the model starts out unable to see far enough back to discover that seeing far back is useful.
Initialising the forget gate bias to 1 or 2 puts it at σ(1)=0.73 or σ(2)=0.88, so the cell begins life leaning towards remembering and learns to forget where forgetting helps.
1for name, param in lstm.named_parameters():2 if "bias" in name:3 n = param.size(0)4 # PyTorch packs gates as [input, forget, cell, output]5 param.data[n // 4: n // 2].fill_(1.0)This is a two-line change that frequently buys a point or more of accuracy on long-sequence tasks, and it is routinely left out.
An LSTM cell from scratch
1import numpy as np23def sigmoid(z):4 return 1.0 / (1.0 + np.exp(-z))56class LSTMCell:7 def __init__(self, input_size, hidden_size, seed=0):8 rng = np.random.default_rng(seed)9 H, Z = hidden_size, hidden_size + input_size10 scale = 1.0 / np.sqrt(Z)11 # One matrix per gate, each acting on [h_prev ; x]12 self.Wf = rng.normal(0, scale, (H, Z))13 self.Wi = rng.normal(0, scale, (H, Z))14 self.Wc = rng.normal(0, scale, (H, Z))15 self.Wo = rng.normal(0, scale, (H, Z))16 self.bf = np.ones((H, 1)) # forget bias = 1, deliberately17 self.bi = np.zeros((H, 1))18 self.bc = np.zeros((H, 1))19 self.bo = np.zeros((H, 1))2021 def step(self, x, h_prev, c_prev):22 z = np.vstack([h_prev, x]) # concatenate once, reuse2324 f = sigmoid(self.Wf @ z + self.bf) # forget25 i = sigmoid(self.Wi @ z + self.bi) # input26 g = np.tanh(self.Wc @ z + self.bc) # candidate27 o = sigmoid(self.Wo @ z + self.bo) # output2829 c = f * c_prev + i * g # the additive path30 h = o * np.tanh(c)31 return h, c, {"f": f, "i": i, "g": g, "o": o}Inspecting the returned gate values on real inputs is the single most useful diagnostic you have. If the forget gates sit near 0 across the board, your cell is amnesiac. If they sit at exactly 1, it never clears and the state saturates.
Counting parameters
Four gates, each with a matrix of shape dh×(dh+dx) and a bias of size dh:
With dx=100, dh=128: 4(128×228+128)=4×29,312=117,248. Exactly four times a plain RNN of the same size, which is the price of the gates.
The GRU: the same idea with less machinery
The Gated Recurrent Unit asks whether four gate computations are necessary. It merges the cell and hidden state into one, and merges the forget and input gates into a single update gate.
The last line is the design in miniature. There is one dial. Whatever fraction of the state you overwrite is exactly the fraction you stop keeping — (1−zt) and zt must sum to 1. An LSTM can independently keep everything and add a lot; a GRU cannot.
The reset gate is the genuinely different piece: it can zero out the previous state before computing the candidate, letting the cell propose new content that ignores history entirely. That is useful at boundaries — the start of a new sentence, a topic shift.
Three quarters of an LSTM's parameters.
| LSTM | GRU | |
|---|---|---|
| Gates | 3 (forget, input, output) | 2 (update, reset) |
| State vectors | 2 (h and c) | 1 (h) |
| Parameters (dx=100, dh=128) | 117,248 | 87,936 |
| Training speed | Baseline | ~25–30% faster per epoch |
| Separate memory from output | Yes — output gate hides state | No — state is the output |
| Very long dependencies (200+ steps) | Usually better | Usually adequate |
| Small datasets | More prone to overfit | Fewer parameters, often better |
| Empirical accuracy on typical NLP | Within noise of each other on most benchmarks |
There is no reliable winner. Start with a GRU because it trains faster, and switch to an LSTM only if your task genuinely involves very long dependencies or the GRU plateaus. Do not spend a week choosing; spend it on your data.
Both in PyTorch
1import torch2import torch.nn as nn34class SequenceClassifier(nn.Module):5 def __init__(self, vocab_size, embed_dim=100, hidden_dim=128,6 num_layers=2, num_classes=2, cell="lstm", dropout=0.3):7 super().__init__()8 self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0)9 rnn_cls = nn.LSTM if cell == "lstm" else nn.GRU10 self.rnn = rnn_cls(11 embed_dim, hidden_dim, num_layers=num_layers,12 batch_first=True,13 dropout=dropout if num_layers > 1 else 0.0, # only between layers14 )15 self.dropout = nn.Dropout(dropout)16 self.fc = nn.Linear(hidden_dim, num_classes)1718 if cell == "lstm":19 for name, p in self.rnn.named_parameters():20 if "bias" in name:21 n = p.size(0)22 p.data[n // 4: n // 2].fill_(1.0)2324 def forward(self, x, lengths):25 emb = self.dropout(self.embedding(x))26 packed = nn.utils.rnn.pack_padded_sequence(27 emb, lengths.cpu(), batch_first=True, enforce_sorted=False28 )29 out, state = self.rnn(packed)30 h_n = state[0] if isinstance(state, tuple) else state # LSTM returns (h, c)31 return self.fc(self.dropout(h_n[-1]))Note that nn.LSTM returns (output, (h_n, c_n)) while nn.GRU returns (output, h_n). Code that assumes one shape and receives the other fails in confusing ways, so handle it explicitly as above.
Adding attention over the sequence
Using only the final hidden state throws away every intermediate state. Attention computes a weighted average of all of them, letting the model choose which positions matter.
1class AttentionPool(nn.Module):2 """Learns a score per timestep, then averages states by softmax weight."""3 def __init__(self, hidden_dim):4 super().__init__()5 self.score = nn.Linear(hidden_dim, 1)67 def forward(self, outputs, mask):8 # outputs: (batch, seq, hidden) mask: (batch, seq) True where real9 scores = self.score(outputs).squeeze(-1) # (batch, seq)10 scores = scores.masked_fill(~mask, float("-inf")) # ignore padding11 weights = torch.softmax(scores, dim=1) # (batch, seq)12 context = torch.bmm(weights.unsqueeze(1), outputs) # (batch, 1, hidden)13 return context.squeeze(1), weightsThe masked_fill line is not optional. Padding positions still produce scores, and without masking they receive a share of the softmax probability, diluting the real content and making attention weights uninterpretable.
The returned weights are directly readable — plot them over the tokens and you can see which words the model used. On the review at the top of this lesson, a working model puts its mass on snapped and will not answer.
What still goes wrong
| Symptom | Cause | Fix |
|---|---|---|
Loss becomes NaN | Exploding gradients — gates do not prevent these | clip_grad_norm_(params, 5.0), always |
| Training accuracy 99%, validation 72% | Overfitting; recurrent models memorise readily | Dropout, pretrained embeddings, early stopping, fewer units |
| Long documents still classified on their endings | Only the final state is used | Attention pooling or max-pooling over all states |
Passing dropout=0.3 with num_layers=1 does nothing | PyTorch applies it only between layers | Apply nn.Dropout yourself on inputs and outputs |
| Very slow training | Recurrence cannot be parallelised across time | Shorter sequences, larger batches, GRU over LSTM |
| Model works on short inputs, degrades on long | Padding processed as real tokens | pack_padded_sequence |
Two of these deserve emphasis because they surprise people.
Gating does not remove the need for gradient clipping. The additive cell path protects against vanishing, not exploding. A single bad batch can still blow up the weights. Clipping costs one line and prevents a class of training failure that otherwise wastes hours.
Sequential computation is the permanent cost. Step t cannot start until step t−1 finishes, so a 200-token sequence requires 200 dependent GPU operations regardless of how much hardware you have. This is not a tuning problem — it is intrinsic to recurrence, and it is the main reason large models moved to architectures that process all positions at once.
When you sit down to build one: start with a single-layer GRU at 128 units on pretrained embeddings, clip at 5.0, pack your sequences, and get a baseline. Then change one thing at a time. Add attention pooling before you add layers; it usually helps more and costs less. And whenever a recurrent model is mysteriously mediocre, print the gate values on a real batch before you touch the hyperparameters — the answer is usually visible there.