Skip to main content

πŸ“ Recurrent Neural Network (RNN)

Description​

< What is it? >​

A Recurrent Neural Network (RNN) is a neural network for ordered data such as text, speech, sensor readings, or time series. It reads one step at a time and carries a hidden state forward, so later steps can use information from earlier ones.

x₁ β†’ [RNN cell] β†’ h₁ β†’ [RNN cell] β†’ hβ‚‚ β†’ [RNN cell] β†’ h₃
↑ ↑ ↑
xβ‚‚ x₃ xβ‚„

The same RNN cell and its weights are reused at every step. This page uses vanilla RNN to mean the simple Elman-style RNN; LSTM is a more capable gated RNN.

Key points​

< Hidden state and recurrence >​

At time step tt, the RNN combines the current input xtx_t with its previous hidden state htβˆ’1h_{t-1}:

ht=tanh⁑(Wxhxt+Whhhtβˆ’1+bh)h_t = \tanh(W_{xh}x_t+W_{hh}h_{t-1}+b_h)

The hidden state is a learned summary of the sequence seen so far. The matrix WxhW_{xh} processes the current input, while WhhW_{hh} carries information from the previous time step. Both matrices are shared across every step of the sequence.

  • Example: while reading β€œThe animal rested because it was tired,” the state at it can contain information about the earlier word animal.

< Unrolling through time >​

An RNN is easiest to visualize as the same cell copied across time:

x₁, hβ‚€ β†’ cell β†’ h₁
xβ‚‚, h₁ β†’ cell β†’ hβ‚‚
x₃, hβ‚‚ β†’ cell β†’ h₃

The copies are not separate models: they all use the same parameters. This weight sharing lets an RNN accept different sequence lengths without adding new parameters for each position.

< Common output patterns >​

PatternInputOutputExample
Many-to-oneA sequenceOne prediction after the final stepPredict sentiment for a review
Many-to-many, alignedA sequenceOne prediction per input stepTag each word for NER
Many-to-many, shiftedA sequenceA next-step prediction at every stepPredict the next token in language modeling

< Training and limitations >​

RNNs are trained with backpropagation through time (BPTT): the network is unrolled over its sequence steps, and gradients flow backward through those steps.

This repeated recurrence is the RNN's main weakness. Gradients can vanish or explode over long sequences, making distant information difficult to learn. RNNs also process steps sequentially, so they cannot parallelize across time as easily as Transformers.

LSTM and GRU cells add gates that make long-range learning more reliable. For modern large-scale language modeling, Transformers have largely replaced vanilla RNNs.

< In PyTorch >​

import torch
import torch.nn as nn

# Input shape: (batch, time_steps, input_size)
x = torch.randn(32, 20, 128)

rnn = nn.RNN(
input_size=128,
hidden_size=256,
batch_first=True,
)

output, h_n = rnn(x)

# output: one hidden state per time step, shape (32, 20, 256)
# h_n: final hidden state, shape (1, 32, 256)

Comparison​

< Vanilla RNN vs LSTM >​

Vanilla RNNLSTM
MemoryOne hidden state hth_tHidden state plus a cell state ctc_t
ControlsNo explicit memory gatesForget, input, and output gates
Long dependenciesOften difficult because gradients vanish or explodeMore reliable, though not perfect
Compute per stepLowerHigher because it computes several gates
Typical use todaySmall or simple sequential baselinesTime series, speech, and legacy sequence models

Reference​