“What is the difference between an RNN, an LSTM and a GRU, and why were gates needed?” still comes up in machine learning interviews, even though transformers dominate most sequence work. It forces you to reason about backpropagation through time, why a product of many small numbers disappears, and the tricks (clipping, truncation, masking) that keep recurrent training stable. Saying “LSTMs have memory” is not enough; writing the cell-state update and explaining why it is additive is.
Before you start
You need the chain rule, the idea that backpropagation multiplies local derivatives along a path, and the sigmoid (0..1) and tanh (-1..1) functions. The examples use Python 3, the standard library and scalar hidden states, so every number can be checked by hand. The PyTorch code in the worked scenario is illustrative for PyTorch 2.x and was not executed here.
The short answer
A vanilla RNN updates a hidden state with the same weights at every step: h_t = tanh(W_hh h_{t-1} + W_xh x_t + b). Training unrolls the loop and backpropagates through time, so the gradient reaching an early step is a product of one Jacobian per step; if those factors are below 1 it vanishes, above 1 it explodes. An LSTM adds a separate cell state updated as c_t = f_t * c_{t-1} + i_t * g_t, where learned gates decide what to forget, what to write and what to output; the additive path lets information and gradients survive many steps. A GRU merges cell and hidden state and uses two gates (update and reset), so it has three weight blocks instead of four and often performs similarly. Clipping controls explosions; gating mitigates vanishing.
How it works
Unroll a scalar RNN for T steps and h_T depends on h_0 through a chain of steps. The local derivative of each link is w * (1 - h_t^2), the recurrent weight times the tanh derivative. Backpropagation through time (BPTT) multiplies those links:
import math
def rnn_forward(xs, w, u, h0=0.0):
"""Scalar RNN: h_t = tanh(w * h_{t-1} + u * x_t). Returns all hidden states."""
hs = [h0]
for x in xs:
hs.append(math.tanh(w * hs[-1] + u * x))
return hs
def grad_h_T_wrt_h0(xs, w, u, h0=0.0):
"""BPTT: multiply the local Jacobians dh_t/dh_{t-1} = w * (1 - h_t^2)."""
hs = rnn_forward(xs, w, u, h0)
g = 1.0
for h in hs[1:]:
g *= w * (1 - h * h)
return gWith vectors, w becomes the matrix W_hh and each link is diag(1 - h_t^2) W_hh. If the largest singular value of every link stays below 1 the gradient shrinks geometrically; if it exceeds 1 the gradient can grow geometrically.
The LSTM changes the structure of the chain. Four blocks of weights read x_t and h_{t-1} and produce a forget gate f, an input gate i, a candidate g and an output gate o. The cell state is updated by addition, and the hidden state is a gated view of it. Along the direct cell path, dc_t / dc_{t-1} is just f_t, a number the network learns and can keep close to 1, instead of a weight matrix times a squashing derivative.
Step-by-step walkthrough
Step 1: Watch the gradient shrink as w to the power T
Without the tanh, the factor is exactly w ** T:
for w in (0.5, 0.9, 1.0, 1.1, 1.5):
print(w, [f"{w ** T:.3g}" for T in (10, 50, 100)])
# 0.5 ['0.000977', '8.88e-16', '7.89e-31']
# 0.9 ['0.349', '0.00515', '2.66e-05']
# 1.0 ['1', '1', '1']
# 1.1 ['2.59', '117', '1.38e+04']
# 1.5 ['57.7', '6.38e+08', '4.07e+17']A weight of 0.9 looks harmless, yet after 100 steps the signal is 0.00003 of its size. A weight of 1.1 multiplies it by about 14,000. Exactly 1 is a knife edge, and real Jacobians vary from step to step. Now include the tanh on random inputs from a seeded generator:
import random
rng = random.Random(0)
xs = [rng.uniform(-1, 1) for _ in range(100)]
for T in (5, 20, 50, 100):
print(T, f"{grad_h_T_wrt_h0(xs[:T], w=0.9, u=0.5):.3e}")
# 5 3.365e-01
# 20 1.839e-03
# 50 2.029e-07
# 100 2.461e-16The tanh derivative is at most 1, so it only makes vanishing worse. By step 100 the first state has no measurable influence, which is why vanilla RNNs struggle to learn that a word fifty tokens back determines the current one.
Step 2: Run one LSTM step on tiny numbers
def sigmoid(z):
return 1 / (1 + math.exp(-z))
def lstm_step(x, h_prev, c_prev, p):
"""One LSTM step with scalar input and hidden size 1. p holds (w_x, w_h, b) per gate."""
pre = {g: p[g][0] * x + p[g][1] * h_prev + p[g][2] for g in "figo"}
f = sigmoid(pre["f"]) # forget gate: how much old cell state to keep
i = sigmoid(pre["i"]) # input gate: how much of the candidate to write
g = math.tanh(pre["g"]) # candidate values
o = sigmoid(pre["o"]) # output gate: how much of the cell to expose
c = f * c_prev + i * g # additive update
h = o * math.tanh(c)
return h, c, {"f": f, "i": i, "g": g, "o": o}
params = {
"f": (0.5, 0.5, 1.0), # forget bias of 1.0 keeps memory by default
"i": (1.0, -0.5, 0.0),
"g": (0.8, 0.2, 0.0),
"o": (0.3, 0.6, 0.0),
}
h, c, gates = lstm_step(x=1.0, h_prev=0.5, c_prev=2.0, p=params)
print({k: round(v, 4) for k, v in gates.items()}) # {'f': 0.852, 'i': 0.6792, 'g': 0.7163, 'o': 0.6457}
print(round(c, 4), round(h, 4)) # 2.1904 0.6297The forget gate keeps 85 percent of the old cell value 2.0 and the input gate writes 68 percent of the candidate 0.716, so the new cell is 0.852 * 2.0 + 0.679 * 0.716 = 2.19. The gates are sigmoids because they act as soft switches between 0 and 1; the candidate is a tanh because it is content that can be positive or negative. Initialising the forget bias to about 1 is a common trick (PyTorch initialises all biases uniformly, so you set it yourself) so the cell remembers by default early in training.
The payoff is the gradient along the cell path: a product of forget-gate values, which the model learns per unit and per step.
for f in (0.5, 0.9, 0.99):
print(f, f"{f ** 100:.3g}")
# 0.5 7.89e-31
# 0.9 2.66e-05
# 0.99 0.366An LSTM does not abolish vanishing gradients. It makes the decay rate a learned, input-dependent quantity, and a unit that needs long memory can hold f near 1.
Step 3: Compare with a GRU step
def gru_step(x, h_prev, p):
z = sigmoid(p["z"][0] * x + p["z"][1] * h_prev + p["z"][2]) # update gate
r = sigmoid(p["r"][0] * x + p["r"][1] * h_prev + p["r"][2]) # reset gate
n = math.tanh(p["n"][0] * x + r * (p["n"][1] * h_prev) + p["n"][2])
return (1 - z) * n + z * h_prev, {"z": z, "r": r, "n": n}
gp = {"z": (0.5, 0.5, 1.0), "r": (1.0, -0.5, 0.0), "n": (0.8, 0.2, 0.0)}
hg, gg = gru_step(1.0, 0.5, gp)
print({k: round(v, 4) for k, v in gg.items()}, round(hg, 4))
# {'z': 0.852, 'r': 0.6792, 'n': 0.7003} 0.5297The GRU has no separate cell. The update gate z interpolates between keeping the old state and taking the candidate n, so one gate plays the role of both forget and input gates. The reset gate r decides how much of the previous state feeds into the candidate. This follows PyTorch’s convention, where z near 1 keeps the old state; some texts swap z and 1 - z, so state the convention you are using.
Step 4: Count parameters
def rnn_params(x, h, gates=1, pytorch=False):
biases = 2 * h if pytorch else h
return gates * (h * x + h * h + biases)
for name, gates in (("RNN", 1), ("GRU", 3), ("LSTM", 4)):
print(name, rnn_params(128, 256, gates), rnn_params(128, 256, gates, pytorch=True))
# RNN 98560 98816
# GRU 295680 296448
# LSTM 394240 395264Each gate is a small layer reading [h_{t-1}, x_t], so it holds h*(h+x) weights and h biases. The LSTM has four, the GRU three, the vanilla RNN one. PyTorch keeps two bias vectors per gate (bias_ih and bias_hh), adding h per gate. For the LSTM the two are mathematically redundant; for the GRU they are not, because b_hn sits inside the reset-gate product.
Worked scenario
A team trains an LSTM on sensor logs of 2,000 steps, backpropagating through each full sequence. Loss falls for a few hundred iterations, then jumps and becomes nan. Logging the gradient norm shows spikes into the thousands just before the failure: an exploding gradient from one long sequence produced an update large enough to push weights into a region where activations overflow, and nan spread through every later step.
The fix has two parts. Clip the global gradient norm between backward() and step(), and train with truncated BPTT, carrying the hidden state forward between 100-step chunks but cutting the graph with detach(). The illustrative PyTorch 2.x loop:
import torch
from torch import nn
import torch.nn.functional as F
lstm = nn.LSTM(input_size=32, hidden_size=256, batch_first=True)
head = nn.Linear(256, 32)
params = list(lstm.parameters()) + list(head.parameters())
opt = torch.optim.Adam(params, lr=1e-3)
for x, y in loader: # x: (batch, 2000, 32), y: (batch, 2000) class ids
state = None
for start in range(0, x.size(1), 100): # truncated BPTT in 100-step chunks
out, state = lstm(x[:, start:start + 100], state)
loss = F.cross_entropy(head(out).flatten(0, 1), y[:, start:start + 100].flatten())
opt.zero_grad()
loss.backward()
grad_norm = nn.utils.clip_grad_norm_(params, max_norm=1.0) # returns the norm before clipping
opt.step()
state = tuple(s.detach() for s in state) # keep the values, drop the graphClipping rescales the whole gradient vector when its norm exceeds max_norm, keeping its direction. The same rule in plain Python:
def clip_by_global_norm(grads, max_norm):
total = math.sqrt(sum(g * g for g in grads))
scale = min(1.0, max_norm / (total + 1e-6))
return [g * scale for g in grads], total
clipped, norm = clip_by_global_norm([30.0, -40.0, 0.0], max_norm=1.0)
print(round(norm, 3), [round(g, 4) for g in clipped]) # 50.0 [0.6, -0.8, 0.0]
print(clip_by_global_norm([0.3, -0.4], max_norm=1.0)) # ([0.3, -0.4], 0.5)Forgetting the detach() gives a different failure: the second chunk’s backward() tries to go through the first chunk’s freed graph and PyTorch raises “Trying to backward through the graph a second time”. Truncation trades away gradients for dependencies longer than the chunk, though the forward pass still carries the state across chunks.
Common mistake
- “LSTMs solve vanishing gradients.” They mitigate it with an additive, gated path; long dependencies are still hard.
- “Clipping fixes vanishing gradients.” Clipping only caps large gradients. It cannot amplify small ones.
- Letting the model read padding. In a padded batch, the final hidden state of a short sequence has run over padding. Use
nn.utils.rnn.pack_padded_sequence(x, lengths, batch_first=True, enforce_sorted=False)(lengths on the CPU) or gather the state at each sequence’s true last step, and mask padded positions out of the loss. - Quoting PyTorch parameter counts as textbook counts. The extra bias vector adds
hper gate. - “GRUs are always better because they are smaller.” They are cheaper; which wins is task dependent.
Verify the behavior
These checks compare the BPTT gradient with a finite difference, confirm that a saturated forget gate preserves the cell over 1,000 steps, and pin the parameter formulas:
def test_bptt_matches_finite_difference():
eps, seq = 1e-6, xs[:20]
numeric = (rnn_forward(seq, 0.9, 0.5, eps)[-1] - rnn_forward(seq, 0.9, 0.5, -eps)[-1]) / (2 * eps)
assert abs(numeric - grad_h_T_wrt_h0(seq, 0.9, 0.5)) < 1e-8
def test_open_forget_gate_keeps_memory():
keep = {"f": (0.0, 0.0, 50.0), "i": (0.0, 0.0, -50.0), "g": (1.0, 0.0, 0.0), "o": (0.0, 0.0, 0.0)}
h, c = 0.0, 2.0
for x in xs * 10:
h, c, _ = lstm_step(x, h, c, keep)
assert abs(c - 2.0) < 1e-12
def test_parameter_counts():
assert rnn_params(128, 256, 4) == 4 * (256 * (256 + 128) + 256)
assert rnn_params(128, 256, 4, pytorch=True) - rnn_params(128, 256, 4) == 4 * 256
test_bptt_matches_finite_difference(); test_open_forget_gate_keeps_memory(); test_parameter_counts()
print("ok")Follow-up questions
Why did transformers replace RNNs for most sequence tasks? An RNN must process step t before t + 1, so training cannot parallelise across time, and information between distant positions travels through many steps. Self-attention connects every pair of positions in one layer and parallelises well on GPUs. The cost is attention compute that grows quadratically with sequence length, and a KV cache that grows with context at inference.
Where do RNNs still make sense? Streaming and on-device workloads such as keyword spotting, or small time-series models, where constant memory per step and low latency matter more than parallel training. Recent state-space models such as Mamba revive the recurrent idea with parallelisable training, but they are a different architecture, not an LSTM variant.
When can you not use a bidirectional RNN? It runs forward and backward passes and concatenates them, so it needs the whole sequence: no streaming, no autoregressive generation.
Interview exercise
You build nn.LSTM(input_size=100, hidden_size=200, num_layers=2) in PyTorch. How many parameters does it have, and how many would the textbook formula give? Then explain what changes if you swap it for a GRU of the same shape, and why.
Answer and reasoning
The trap is that layer 2 reads layer 1’s hidden state, so its input size is 200, not 100. With PyTorch’s two biases per gate, layer 1 has 4 * (200*100 + 200*200 + 2*200) = 241,600 and layer 2 has 4 * (200*200 + 200*200 + 2*200) = 321,600, for 563,200 in total. The textbook formula, 4 * (h*(h+x) + h) per layer, gives 240,800 and 320,800, so 561,600; the gap of 1,600 is exactly one extra 200-length bias for each of four gates in two layers. A GRU of the same shape has three gate blocks instead of four, so every count is multiplied by three quarters: 422,400 in PyTorch. It is cheaper per step and often matches the LSTM, but only a validation comparison on your task decides.