Recurrent neural networks
Recurrent neural networks (RNNs) are sequential models that remember information from earlier elements in a sequence.
A sequence is a set of samples where the order matters, like the words in a sentence.
In sequence processing, each element in the input is called a token (often a word).
Algorithms that can understand sequences can also generate new ones, including stories, music, or lyrics.
Working with language
Common natural language processing tasks
NLP (natural language processing) is the processing of natural language information by a computer.
Common NLP tasks include:
- sentiment analysis
- translation
- question answering
- summarization / paraphrasing
- text generation:
- the system starts with a prompt (or seed)
- it predicts the next word based on that prompt
- it appends the predicted word to the prompt
- the updated prompt becomes the input to predict the next word
- this technique is called autoregression:
- regression → a technique for estimating or predicting an output value based on input variables
- auto → using the sequence's own previous outputs as inputs
- autoregressive systems are called autoregressors
- the broader term for algorithmic text creation is natural language generation (NLG)
- logical flow (checking whether conclusions follow from premises):
- especially challenging
- often need human–computer collaboration
These tasks rely on language models, which don't understand language or the meaning of words. Instead, they create outputs that seem correct by relying purely on statistical methods.
Transforming text into numbers
Written text must first be converted into numbers.
Two main ways to convert text into numbers:
- character-based encoding: each character gets a unique number
- word-based encoding: each word gets a unique number
Then numbers are fed into an autoregressive neural network, which predicts the next number in the sequence and then turns those numbers back into words to see the generated text.
Fine-tuning and downstream networks
Models can start as general-purpose systems and be adapted for specialized tasks.
This specialization leverages the knowledge the original model has already learned.
Ways to specialize a system:
- transfer learning (add layers)
- freeze a pretrained model
- add new layers at the end of the classification section
- train the new layers on specialized data
- the pretrained model's knowledge is "transferred" to the specialized task
- fine-tuning (full model adjustment)
- start with a pretrained language model (trained on a general dataset)
- train it further on specialized data
- all weights are adjusted to specialize for the new domain
- downstream network
- freeze a pretrained model
- feed its output into a second model (a downstream network) designed for a specific task
In practice, many systems combine fine-tuning and downstream networks.
Neural networks fail at language prediction
A tiny fully connected neural network is used to predict sequences:
The network is trained to make predictions using the sliding window approach:
- choose a window size (e.g., 5 items)
- take the first 5 items as input
- use the next as the target/label
- train the network to map the input window to the target
- slide the window forward and repeat
- this teaches a model to predict the next element from a fixed number of previous ones
After training it, this tiny fully connected neural network fails to predict words accurately. Such a tiny network cannot handle complex data like language.
Even a larger fully connected neural network can't handle the complexities of natural language because:
- it can't capture the structure of the text (its semantics)
- predicting the next word requires understanding context from many words earlier
- without this, too many possible continuations remain plausible
- tiny errors in predictions lead to incomprehensible text
- words are stored as single numbers
- tiny differences in predictions can give unrelated words
- example: words "keep" and "flint" are assigned consecutive numbers:
1003and1004- the network predicts the next word as
1003.49 - it is rounded to the nearest integer to convert to a word →
1003→ "keep" - if it predicts 1003.51 instead, rounding gives
1004→ "flint"
- the network predicts the next word as
- it doesn't track the order of words in the input
- this makes it impossible to resolve references (like pronouns) that depend on sequence order
- e.g., "Bob told John that he was hungry" → who does "he" refer to?
We need something smarter than fully connected layers and representing words as single numbers.
Recurrent neural networks
Recurrent neural networks (RNNs) are designed to manage language as an ordered sequence.
They involve new concepts.
State
State is the condition of a system at a given moment which includes:
- the system's current situation
- any information it needs to remember from earlier inputs (possibly in a compressed or transformed form)
State evolves over time. With each new input, the system:
- updates its state
- produces an output based on both the new input and its current state (because this internal state is not visible to observers, it's called the hidden state)
The order of inputs matters: changing the sequence leads to different states and different outputs.
Each input in a sequence is called a time step.
Recurrent cells
Sequential data is processed using recurrent cells, self-contained modules that both compute outputs and manage state over time:
- input arrives
- hidden state is pulled out of the delay (memory)
- system computes a new state based on the input and its hidden state:
- output is produced
- state is updated and pushed in the delay
- repeat for the next input, which is why it's called recurrent
The system contains one or multiple neural networks that learn how to manage state and produce outputs during training.
The cell's internal state:
- is stored as a tensor (often a one-dimensional list of numbers), whose length is called the width or size of the cell
- is usually private but can be exported because some networks can make good use of this information
A recurrent cell on its own layer forms a recurrent layer, and networks built mainly from these are RNNs. Depending on context, "RNN" may refer to the network, the layer, or the cell itself.
Diagrams of long sequences can be unrolled (showing each step separately) or rolled-up (a compact version).
Recurrent cells in action
An unrolled RNN diagram illustrates how a recurrent cell predicts the next word in a five-word sequence:
- the hidden state starts as all zeros
- inputs enter from the bottom
- predictions exit from the top
- each word of the five-word sequence is fed in sequentially
- at each step, the cell:
- uses the current input and state to produce a next-word prediction
- uses the information it learned during training to update its hidden state to encode the context seen so far
As the sequence progresses, the hidden state accumulates a compact representation of the words seen:
- first
it+ predictionswam - then
it was+ predictionnight - then
it was the+ predictionbest - and so on…
Early in training, predictions may be inaccurate, but after enough training on real text, the RNN learns to represent the sequence effectively in its hidden state and assign high probability to the correct next word.
Backpropagation through time
Even though the above diagram shows several "cells", it is really the same recurrent cell reused at each time step with shared weights.
To train that recurrent cell:
- gradients must be computed starting from the last time step
- then propagated backward through earlier time steps
To account for this, the gradient of the final error must be propagated backward through all previous steps to update the shared weights.
The solution is backpropagation through time (BPTT):
- a special variant of backpropagation used to train RNNs
- applies backpropagation through the sequence of time steps
- handles the gradient computations required to train a recurrent cell with shared weights
However, backpropagating gradients through many time steps can cause fundamental challenges in training RNNs:
- vanishing gradients: gradients become smaller and smaller, slowing or stopping learning
- exploding gradients: gradients become larger and larger, destabilizing training
Long short-term memory and gated recurrent networks
A long short-term memory (LSTM) network is a type of recurrent cell designed to prevent vanishing and exploding gradients.
It maintains an internal state (memory) that can be selectively updated over time:
- the state changes frequently, acting like short-term memory
- some information can be kept for a long time
- think of it as a selectively persistent short-term memory
An LSTM uses three internal networks that implement gates, collectively known as the gating mechanism:
- forget network: decides which information to remove from the LSTM's state
- remember network: decides what new information to add to the state
- select network: determines what part of the internal state becomes the cell's output
"Forgetting" means pushing values stored in the memory state toward zero, while "remembering" means adding new values to the memory state.
The core difference with basic RNN is the cell state:
- it is updated additively rather than purely multiplicatively
- so gradients can flow through many time steps without shrinking or exploding
- and standard backpropagation can be used
LSTMs are so common that "RNN" often means an "LSTM network".
A common variation of the LSTM is the gated recurrent unit (GRU), and both LSTM and GRU are often tested to see which works best.
Different architectures
Recurrent cells can be either stacked or combined with other networks to handle complex sequence tasks.
CNN-LSTM Networks:
- combines convolutional layers with LSTM cells
- especially useful for video classification tasks
- the CNN part identifies objects in input data (such as video frames)
- the LSTM part tracks how those objects move from one frame to the next
Deep RNNs:
- stack multiple RNN or LSTM layers
- each layer's output feeds the next
- layers can specialize in subtasks (e.g., turning text into an abstract form, changing tone, translating)
- each layer can be trained or improved independently, but replacing one may need some retraining
Bidirectional RNNs (bi-RNNs):
- look at the full context of a sentence: both before and after each word
- handles linguistic challenges:
- ambiguity: sentences may have multiple meanings depending on word order, context, or emphasis
- polysemy: words like "cast" have multiple meanings, context before and after is needed for correct translation
- two RNNs run simultaneously: one forward, one backward
- outputs are combined
- multiple bi-RNNs can be stacked to form a deep bi-RNN for more powerful language modeling
Seq2Seq
Seq2Seq ("sequence to sequence") is an algorithm for translating entire sentences between languages, rather than word by word.
Translation challenges:
- different languages use different word orders
- sentence lengths vary across languages
Seq2Seq is conceptually similar to autoencoders, but the "latent vector" is called the context vector:
It uses two RNNs:
- encoder:
- reads the input sentence word by word
- updates its hidden state at each step
- ignores outputs at each step; only the final hidden state matters
- the final hidden state is the context vector, summarizing the input sentence
- decoder:
- starts with the context vector as its initial hidden state
- starts the output sequence with a
[START]token - generates a sequence autoregressively: each generated word becomes the next input
- stops when it produces an
[END]token
This approach allows Seq2Seq to generate variable-length output sequences from fixed-length input sequences.
Limitations of Seq2Seq:
- context vector bottleneck: all input information must fit into a single fixed-size context vector, which can be insufficient for long sentences
- long-term dependency problem: long sentences can exceed the encoder's memory capacity and be forgotten by the time the final hidden state is produced
- sequential processing: Seq2Seq models are built on RNNs which are inherently sequential, so training and generation cannot be parallelized
Seq2Seq is simple, widely used, and easy to implement. It works well for short sentences but struggles with long or complex ones.