Natural Language Processing Basics

LSTMs & GRUs — Solving the Vanishing Gradient Problem


Train a plain recurrent network on product reviews and then feed it this one:

Text
"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.

The gates decide what the cell keepsForget gate:how much ofc to eraseInput gate:how muchnew to admitCandidate:what the newcontent isNew cell:forget timesc, plus newOutput gate:what toreveal as h_tSet the forget-gate bias to 1 and the cell starts out remembering rather than erasing.
The cell update is addition, not repeated multiplication, and that single change is what lets gradient travel far.

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:

ht=tanh⁡(Wxhxt+Whhht−1+b)h_t = \tanh(W_{xh} x_t + W_{hh} h_{t-1} + b)

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:

∂ht∂ht−1=Whh⊤ diag(1−ht2)\frac{\partial h_t}{\partial h_{t-1}} = W_{hh}^\top \, \text{diag}(1 - h_t^2)

Two things guarantee that this shrinks. The tanh⁡\tanh derivative 1−ht21 - h_t^2 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−60t-60 is 0.760≈5×10−100.7^{60} \approx 5 \times 10^{-10} of the gradient at step tt.

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 ctc_t — the long-term memory. It is modified by addition, not by being overwritten.
  • Hidden state hth_t — the working output, a filtered view of ctc_t, 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)(0,1):

ft=σ(Wf[ht−1,xt]+bf)forget gatef_t = \sigma(W_f [h_{t-1}, x_t] + b_f) \qquad \textbf{forget gate}

"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.

it=σ(Wi[ht−1,xt]+bi)input gatei_t = \sigma(W_i [h_{t-1}, x_t] + b_i) \qquad \textbf{input gate}

"How much of the newly proposed information should I write in?"

c~t=tanh⁡(Wc[ht−1,xt]+bc)candidate\tilde{c}_t = \tanh(W_c [h_{t-1}, x_t] + b_c) \qquad \textbf{candidate}

The proposed new content, in (−1,1)(-1, 1). Note this uses tanh⁡\tanh, not sigmoid — it is content, not a gate.

ct=ft⊙ct−1+it⊙c~tmemory updatec_t = f_t \odot c_{t-1} + i_t \odot \tilde{c}_t \qquad \textbf{memory update}

This is the line that matters. Old memory, scaled by the forget gate, plus new content, scaled by the input gate. ⊙\odot is element-wise multiplication.

ot=σ(Wo[ht−1,xt]+bo)output gateo_t = \sigma(W_o [h_{t-1}, x_t] + b_o) \qquad \textbf{output gate}

ht=ot⊙tanh⁡(ct)what gets exposedh_t = o_t \odot \tanh(c_t) \qquad \textbf{what gets exposed}

The memory is not the output. The cell can hold something in ctc_t for fifty steps while keeping oto_t 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:

∂ct∂ct−1=ft\frac{\partial c_t}{\partial c_{t-1}} = f_t

Compare this with the plain RNN's Whh⊤diag(1−ht2)W_{hh}^\top \text{diag}(1 - h_t^2). There is no weight matrix. There is no tanh⁡\tanh derivative. The gradient along the cell state is multiplied by the forget gate and nothing else.

If the network learns ft≈1f_t \approx 1 for the dimensions holding something important, the gradient passes through essentially unchanged, however many steps it travels. Put numbers on it:

Steps backPlain RNN, factor 0.7LSTM, ft=0.99f_t = 0.99LSTM, ft=0.9f_t = 0.9
100.0280.9040.349
501.8×10−81.8 \times 10^{-8}0.6050.005
1003.2×10−163.2 \times 10^{-16}0.3662.7×10−52.7 \times 10^{-5}
2001.0×10−311.0 \times 10^{-31}0.1347.1×10−107.1 \times 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 ctc_t back to ct−1c_{t-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.

Text
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.512

On 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\sigma(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\sigma(1) = 0.73 or σ(2)=0.88\sigma(2) = 0.88, so the cell begins life leaning towards remembering and learns to forget where forgetting helps.

Python
for name, param in lstm.named_parameters():    if "bias" in name:        n = param.size(0)        # PyTorch packs gates as [input, forget, cell, output]        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

Python
import numpy as npdef sigmoid(z):    return 1.0 / (1.0 + np.exp(-z))class LSTMCell:    def __init__(self, input_size, hidden_size, seed=0):        rng = np.random.default_rng(seed)        H, Z = hidden_size, hidden_size + input_size        scale = 1.0 / np.sqrt(Z)        # One matrix per gate, each acting on [h_prev ; x]        self.Wf = rng.normal(0, scale, (H, Z))        self.Wi = rng.normal(0, scale, (H, Z))        self.Wc = rng.normal(0, scale, (H, Z))        self.Wo = rng.normal(0, scale, (H, Z))        self.bf = np.ones((H, 1))      # forget bias = 1, deliberately        self.bi = np.zeros((H, 1))        self.bc = np.zeros((H, 1))        self.bo = np.zeros((H, 1))    def step(self, x, h_prev, c_prev):        z = np.vstack([h_prev, x])            # concatenate once, reuse        f = sigmoid(self.Wf @ z + self.bf)    # forget        i = sigmoid(self.Wi @ z + self.bi)    # input        g = np.tanh(self.Wc @ z + self.bc)    # candidate        o = sigmoid(self.Wo @ z + self.bo)    # output        c = f * c_prev + i * g                # the additive path        h = o * np.tanh(c)        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)d_h \times (d_h + d_x) and a bias of size dhd_h:

∣θ∣=4(dh(dh+dx)+dh)|\theta| = 4\left(d_h (d_h + d_x) + d_h\right)

With dx=100d_x = 100, dh=128d_h = 128: 4(128×228+128)=4×29,312=117,2484(128 \times 228 + 128) = 4 \times 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.

zt=σ(Wz[ht−1,xt])update gatez_t = \sigma(W_z [h_{t-1}, x_t]) \qquad \textbf{update gate}

rt=σ(Wr[ht−1,xt])reset gater_t = \sigma(W_r [h_{t-1}, x_t]) \qquad \textbf{reset gate}
h~t=tanh⁡(Wh[rt⊙ht−1,xt])candidate\tilde{h}_t = \tanh(W_h [r_t \odot h_{t-1}, x_t]) \qquad \textbf{candidate}
ht=(1−zt)⊙ht−1+zt⊙h~th_t = (1 - z_t) \odot h_{t-1} + z_t \odot \tilde{h}_t

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)(1 - z_t) and ztz_t 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.

∣θ∣GRU=3(dh(dh+dx)+dh)|\theta|_{\text{GRU}} = 3\left(d_h(d_h + d_x) + d_h\right)

Three quarters of an LSTM's parameters.

LSTMGRU
Gates3 (forget, input, output)2 (update, reset)
State vectors2 (hh and cc)1 (hh)
Parameters (dx=100d_x{=}100, dh=128d_h{=}128)117,24887,936
Training speedBaseline~25–30% faster per epoch
Separate memory from outputYes — output gate hides stateNo — state is the output
Very long dependencies (200+ steps)Usually betterUsually adequate
Small datasetsMore prone to overfitFewer parameters, often better
Empirical accuracy on typical NLPWithin 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

Python
import torchimport torch.nn as nnclass SequenceClassifier(nn.Module):    def __init__(self, vocab_size, embed_dim=100, hidden_dim=128,                 num_layers=2, num_classes=2, cell="lstm", dropout=0.3):        super().__init__()        self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0)        rnn_cls = nn.LSTM if cell == "lstm" else nn.GRU        self.rnn = rnn_cls(            embed_dim, hidden_dim, num_layers=num_layers,            batch_first=True,            dropout=dropout if num_layers > 1 else 0.0,   # only between layers        )        self.dropout = nn.Dropout(dropout)        self.fc = nn.Linear(hidden_dim, num_classes)        if cell == "lstm":            for name, p in self.rnn.named_parameters():                if "bias" in name:                    n = p.size(0)                    p.data[n // 4: n // 2].fill_(1.0)    def forward(self, x, lengths):        emb = self.dropout(self.embedding(x))        packed = nn.utils.rnn.pack_padded_sequence(            emb, lengths.cpu(), batch_first=True, enforce_sorted=False        )        out, state = self.rnn(packed)        h_n = state[0] if isinstance(state, tuple) else state   # LSTM returns (h, c)        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.

Python
class AttentionPool(nn.Module):    """Learns a score per timestep, then averages states by softmax weight."""    def __init__(self, hidden_dim):        super().__init__()        self.score = nn.Linear(hidden_dim, 1)    def forward(self, outputs, mask):        # outputs: (batch, seq, hidden)   mask: (batch, seq) True where real        scores = self.score(outputs).squeeze(-1)              # (batch, seq)        scores = scores.masked_fill(~mask, float("-inf"))     # ignore padding        weights = torch.softmax(scores, dim=1)                # (batch, seq)        context = torch.bmm(weights.unsqueeze(1), outputs)    # (batch, 1, hidden)        return context.squeeze(1), weights

The 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

SymptomCauseFix
Loss becomes NaNExploding gradients — gates do not prevent theseclip_grad_norm_(params, 5.0), always
Training accuracy 99%, validation 72%Overfitting; recurrent models memorise readilyDropout, pretrained embeddings, early stopping, fewer units
Long documents still classified on their endingsOnly the final state is usedAttention pooling or max-pooling over all states
Passing dropout=0.3 with num_layers=1 does nothingPyTorch applies it only between layersApply nn.Dropout yourself on inputs and outputs
Very slow trainingRecurrence cannot be parallelised across timeShorter sequences, larger batches, GRU over LSTM
Model works on short inputs, degrades on longPadding processed as real tokenspack_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 tt cannot start until step t−1t-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.