Recurrent networks process a sequence one step at a time and carry a fixed-size state from step to step. Attention-based models replaced them for most large-scale language work, but recurrent cells remain the right tool in plenty of places. They suit streaming inference with constant memory per step, small on-device models, sensor and control loops, and as the conceptual ancestor of today's state-space models. LSTM and GRU are also where most engineers first meet gating, and gating now appears everywhere.

This article defines the plain RNN, the LSTM and the GRU exactly, with PyTorch's equations and weight layout. It builds an LSTM forward and backward pass in NumPy and verifies it against finite differences, counts parameters, and then covers the training mechanics that cause most real bugs: truncated backpropagation through time, hidden-state detaching, and padding versus packing.

The plain RNN and why its gradients fail

A plain (Elman) RNN computes ht = tanh(W xt + U ht-1 + b) and reads predictions off ht. The same W, U and b are used at every step, which is what lets one model handle any sequence length. ht is a lossy summary of everything seen so far.

Training unrolls the loop and backpropagates through it (backpropagation through time, BPTT). The gradient of a loss at step t with respect to the state k steps earlier is a product of k Jacobians, each equal to diag(tanh') times U. If the largest singular value of those Jacobians stays below 1, the product shrinks geometrically. A scalar caricature: with U = 0.5 and tanh' at most 1, twenty steps scale the gradient by at most 0.520, about 9.5e-7. Above 1 the product can explode instead. Clipping fixes explosion. Vanishing needs an architectural fix, which is what gates provide.

The LSTM cell, exactly

The LSTM (Hochreiter and Schmidhuber, 1997; the forget gate was added by Gers et al., 2000) adds a cell state c that is updated additively. PyTorch documents it as:

i_t = sigmoid(W_ii x_t + b_ii + W_hi h_{t-1} + b_hi)     # input gate
f_t = sigmoid(W_if x_t + b_if + W_hf h_{t-1} + b_hf)     # forget gate
g_t = tanh   (W_ig x_t + b_ig + W_hg h_{t-1} + b_hg)     # candidate
o_t = sigmoid(W_io x_t + b_io + W_ho h_{t-1} + b_ho)     # output gate
c_t = f_t * c_t-1 + i_t * g_t
h_t = o_t * tanh(c_t)

The key line is ct = ft * ct-1 + it * gt. Along the cell path, the gradient from ct to ct-1 is multiplied by ft elementwise, not by a full weight matrix and a squashing derivative. If the network learns f near 1 for a unit, that unit can carry information and gradient across long gaps. Worked numbers: over 50 steps, a forget gate of 0.9 keeps 0.950 = 0.005 of the signal, while 0.99 keeps 0.605. The difference between remembering and forgetting is a few hundredths in a gate value.

This is why initialisation matters. With the forget bias at 0, f starts near sigmoid(0) = 0.5, and 0.550 is about 8.9e-16, so nothing long-range is learnable early on. Initialising the forget bias to 1 (recommended by Jozefowicz et al., 2015) starts f near sigmoid(1) = 0.73. sigmoid(1)50 is still only about 1.6e-7, so the bias gives early training a gentler start, and the network still has to learn f close to 1 for the units that need long memory.

The GRU, and two conventions that bite

The GRU (Cho et al., 2014) merges cell and output into one state and uses three gates. PyTorch's form:

r_t = sigmoid(W_ir x_t + b_ir + W_hr h_{t-1} + b_hr)         # reset
z_t = sigmoid(W_iz x_t + b_iz + W_hz h_{t-1} + b_hz)         # update
n_t = tanh(W_in x_t + b_in + r_t * (W_hn h_{t-1} + b_hn))    # new / candidate
h_t = (1 - z_t) * n_t + z_t * h_{t-1}

Two conventions trip people up when porting weights or reading papers. First, in PyTorch z near 1 keeps the old state. Some papers and frameworks write the interpolation the other way round, with z weighting the candidate. Second, PyTorch applies the reset gate after the hidden matmul, r * (Whnh + bhn). Cho et al.'s original applies it before, W(r * h). The two are not numerically equivalent, so weights trained in one convention give wrong outputs in the other even though every shape matches.

The GRU's interpolation gives the same additive shortcut as the LSTM cell. When z is near 1, ht is almost ht-1 and the gradient passes through scaled by z.

Data flow and weight layout

One LSTM step: four gates from one matmul, one additive memory pathx_tinput vectorh_t-1previous outputW x + U h + bone fused matmul, 4H wideisigmoidfsigmoidgtanhosigmoidc_t-1cell statec_t = f * c_t-1 + i * gadditive updategradient path: multiplied by f onlyh_t = o * tanh(c_t)exposed outputGRU: same idea with 3 gates (r, z, n) and no separate cell; h itself is the memory.
The four LSTM gates come from one fused matrix multiply of width 4H. The purple path is the cell state, where the step-to-step gradient is an elementwise product with f rather than a dense Jacobian.

PyTorch stores the four gate blocks stacked in the order (i, f, g, o) in weight_ih_l0, shape (4H, input_size), and weight_hh_l0, shape (4H, H). The GRU uses order (r, z, n). To set the forget bias, write to rows H to 2H of bias_ih_l0 or bias_hh_l0. Both are added, and PyTorch initialises both uniformly in plus or minus 1/sqrt(H), not at zero, so set one to 1 and explicitly zero the other.

A gradient-checked LSTM in NumPy

Writing the backward pass once is the best way to understand the cell. The code below uses one fused matrix of width 4H in the same gate order, caches each step, and walks backward accumulating dh (from the step above and the output) and dc (along the cell path):

import numpy as np
sig = lambda x: 1 / (1 + np.exp(-x))

def lstm_grad(W, U, b, xs, y):
    H = U.shape[1]; h = np.zeros(H); c = np.zeros(H); cache = []
    for x in xs:                                   # forward
        z = W @ x + U @ h + b
        i, f, g, o = sig(z[:H]), sig(z[H:2*H]), np.tanh(z[2*H:3*H]), sig(z[3*H:])
        cp, hp = c, h
        c = f * c + i * g; tc = np.tanh(c); h = o * tc
        cache.append((x, hp, cp, i, f, g, o, tc))
    loss = 0.5 * np.sum((h - y) ** 2)              # loss on the final state
    dW, dU, db = np.zeros_like(W), np.zeros_like(U), np.zeros_like(b)
    dh, dc = h - y, np.zeros(H)
    for x, hp, cp, i, f, g, o, tc in reversed(cache):
        do = dh * tc
        dc = dc + dh * o * (1 - tc ** 2)           # through h = o * tanh(c)
        di, df, dg = dc * g, dc * cp, dc * i
        dz = np.concatenate([di * i * (1 - i), df * f * (1 - f),
                             dg * (1 - g ** 2), do * o * (1 - o)])
        dW += np.outer(dz, x); dU += np.outer(dz, hp); db += dz
        dh = U.T @ dz                              # to h_{t-1}
        dc = dc * f                                # to c_{t-1}: the additive path
    return loss, dW, dU, db

We checked every parameter's gradient against central finite differences (epsilon = 1e-5) on a 6-step sequence with input size 3 and hidden size 4. The worst relative error was 2.4e-8. Keep a check like this in your test suite whenever you write a custom cell. A wrong gate slice still trains, just badly, and nothing else will flag it.

Counting parameters and compute

Per layer, with input size x and hidden size h, PyTorch's two bias vectors give:

CellParametersx = 128, h = 256
RNN (tanh)h(x + h) + 2h98,816
GRU3h(x + h) + 6h296,448
LSTM4h(x + h) + 8h395,264

So a GRU is three quarters the size and compute of an LSTM at equal width. Each step costs about 2 x params FLOPs per sequence element, and the recurrent matmul cannot be parallelised across time. That sequential dependency, not the FLOP count, is why RNN training is slow on GPUs relative to attention. Stacked layers use x = h (or 2h if bidirectional) from layer 2 onward.

Truncated BPTT in PyTorch

Long sequences (a continuous sensor stream, a book-length character corpus) cannot be unrolled end to end. Truncated BPTT splits them into chunks of k steps. It carries the hidden state forward across chunks, so the model sees long context in the forward pass, and cuts the graph between chunks, so gradients flow at most k steps:

import torch
from torch import nn

model = nn.LSTM(input_size=32, hidden_size=256, num_layers=2, batch_first=True, dropout=0.1)
head = nn.Linear(256, 32)
opt = torch.optim.Adam(list(model.parameters()) + list(head.parameters()), lr=1e-3)
with torch.no_grad():                         # forget-gate bias = 1 (rows H:2H)
    for name, p in model.named_parameters():
        if name.startswith("bias_ih"):
            p[256:512].fill_(1.0)
        elif name.startswith("bias_hh"):
            p[256:512].zero_()

state = None
for x, y in chunks:                           # x: (batch, k, 32), consecutive in time
    if state is not None:
        state = tuple(s.detach() for s in state)   # cut the graph, keep the values
    out, state = model(x, state)
    loss = nn.functional.mse_loss(head(out), y)
    opt.zero_grad()
    loss.backward()
    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
    opt.step()

Forgetting the detach() is the classic bug. The second chunk's backward pass then reaches into the first chunk's already-freed graph and raises "Trying to backward through the graph a second time". Silencing that with retain_graph=True is worse: the graph then spans every chunk since the start, and memory and step time grow without bound. The opposite bug is resetting state = None every chunk when chunks are consecutive, which throws away the context you meant to carry. Batches must also be laid out so that row b of chunk j+1 continues row b of chunk j.

Padding, packing and final states

Batches of variable-length sequences are padded to a common length. Without care, the RNN runs over the padding and the final hidden state reflects pad tokens, not the last real element. Packing fixes this:

from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence

packed = pack_padded_sequence(x_padded, lengths.cpu(), batch_first=True, enforce_sorted=False)
packed_out, (h_n, c_n) = model(packed)
out, _ = pad_packed_sequence(packed_out, batch_first=True)
summary = h_n[-1]          # last layer, state at each sequence's true final step

With a packed input, h_n holds the state at each sequence's true end. If you instead take out[:, -1] from a padded run, short sequences give you a state after the padding was processed. For bidirectional models, h_n has one entry per direction per layer. The backward direction's final state summarises the sequence from its end to its start, so concatenate both rather than taking the last row only. Losses on per-step outputs must also be masked so padding positions contribute nothing.

Operational guidance

  • Streaming inference. Keep (h, c) per session and feed one step at a time. Memory per stream is constant, unlike a transformer's growing KV cache. Reset state explicitly at session boundaries, and persist it if sessions can move between servers.
  • Training and inference must agree. A model trained on chunks that started from zero state behaves differently when served with carried state, and vice versa. Train the way you serve.
  • Fast kernels. On GPU, PyTorch dispatches nn.LSTM and nn.GRU to cuDNN fused kernels. A hand-written cell in a Python loop is usually many times slower; reserve it for research cells, or compile it.
  • Gradient clipping is not optional. Clip by global norm (values around 0.5 to 5 are common) and log the pre-clip norm; a sudden rise is an early warning of divergence.
  • Regularisation. dropout in nn.LSTM applies between stacked layers only, never on the recurrent connection, and it does nothing with one layer.

Failure modes

  • Undetached state across batches. An immediate error about backward through the graph a second time; if someone adds retain_graph=True to silence it, memory climbs every step.
  • Ported weights in the wrong gate order. Keras, ONNX exports and custom code may order gates differently from PyTorch's (i, f, g, o) and (r, z, n). The model loads, then outputs nonsense. Compare a step's output against a reference on a fixed input.
  • Padding contamination. Unpacked padded batches give final states and losses that depend on batch composition. Test by checking that a sequence's output is unchanged when it is batched with longer ones.
  • Bidirectional in a streaming product. The backward direction needs the future. A model that was validated offline cannot be served causally.
  • Saturated gates. Huge input scales push sigmoids to 0 or 1 and freeze learning. Normalise inputs, and consider layer normalisation inside the cell for deep stacks.

Trade-offs

ModelTrain parallel over timeInference memory per stepLong-range recallUse when
Plain RNNNoO(h)PoorTeaching, tiny control loops
GRUNoO(h)Good up to hundreds of stepsSmall budgets, similar accuracy to LSTM on many tasks
LSTMNoO(h) plus cellGood up to hundreds of stepsStreaming, on-device, well-understood baseline
TransformerYesGrows with context (KV cache)Direct access within the windowLarge data, long exact recall
State-space (Mamba-style)Yes (scan)O(state)Strong, task-dependentLong sequences with streaming needs

Go deeper with the vanishing gradient problem, gradient clipping, the attention mechanism that grew out of RNN encoder-decoders, and state-space model maths, which revives recurrence in a form that trains in parallel.

What to do next

  1. Implement one LSTM step in NumPy and gradient-check it; keep the check as a unit test.
  2. In PyTorch, set the forget-gate bias to 1 and verify by printing rows H to 2H of the bias.
  3. Decide between full BPTT and truncated BPTT from your sequence lengths, and if truncating, detach state between chunks and lay batches out contiguously in time.
  4. Pack variable-length batches, take final states from h_n, and mask per-step losses.
  5. Clip gradients by global norm and log the pre-clip norm per step.
  6. Try a GRU at the same width and an LSTM, and compare validation loss per parameter rather than per layer.
  7. If you need recall over thousands of steps with full parallel training, benchmark a transformer or a state-space model against your best recurrent baseline.
Key takeaway: RNNs share one set of weights across time and carry a fixed-size state. Plain RNNs lose gradients geometrically, and LSTM and GRU fix that with gated, additive state updates, so memory length is governed by how close the gates sit to 1. Most real bugs are mechanical rather than mathematical: wrong gate order, undetached state, padding contamination and unclipped gradients. Recurrence remains the right choice where constant-memory streaming matters more than parallel training.