In 1997 Sepp Hochreiter and Jรผrgen Schmidhuber published the Long Short-Term Memory network, designed specifically to overcome the vanishing-gradient problem they had analysed. For nearly two decades, LSTMs were the dominant architecture for speech recognition, machine translation, handwriting recognition and language modelling. Understanding them teaches the powerful idea of gating โ learned, multiplicative control over information flow โ which reappears in GRUs, gated linear units and modern state-space models.
The key idea: a protected memory lane#
The LSTM maintains two states:
- the cell state $\mathbf{c}_t$ โ a long-term memory "conveyor belt" updated mostly by addition;
- the hidden state $\mathbf{h}_t$ โ the short-term output exposed to the rest of the network.
Three gates, each a sigmoid layer producing values in $(0, 1)$, control the cell.
The equations#
Given input $\mathbf{x}_t$ and previous hidden state $\mathbf{h}_{t-1}$:
Reading the gates#
- Forget gate $\mathbf{f}_t$: what fraction of each memory component to keep. In a language model, it might clear the stored grammatical number of the subject when a new sentence begins.
- Input gate $\mathbf{i}_t$: how much of the new candidate to write.
- Candidate $\tilde{\mathbf{c}}_t$: the new content that could be stored.
- Output gate $\mathbf{o}_t$: which parts of the memory to reveal as the hidden state now.
Why LSTMs preserve gradients#
Look at the cell update: $\mathbf{c}_t = \mathbf{f}_t \odot \mathbf{c}_{t-1} + \dots$. Its Jacobian with respect to the previous cell (along the direct path) is
No repeated multiplication by a weight matrix and no squashing derivative โ just the forget gate. When the network learns to keep $\mathbf{f}_t \approx 1$ for a component, the gradient flows back through many steps almost unchanged. Hochreiter and Schmidhuber called this the constant error carousel. It is the same idea as a residual connection, applied through time โ two decades before ResNets.
Parameters and cost#
Each of the four transformations maps $[\mathbf{h}_{t-1}, \mathbf{x}_t]$ to the hidden size $H$, so an LSTM has
parameters for input size $D$ โ four times a vanilla RNN.
Using LSTMs in PyTorch#
import torch
import torch.nn as nn
class SentimentLSTM(nn.Module):
def __init__(self, vocab, emb=100, hidden=128, layers=2, classes=2):
super().__init__()
self.emb = nn.Embedding(vocab, emb, padding_idx=0)
self.lstm = nn.LSTM(emb, hidden, num_layers=layers, batch_first=True,
bidirectional=True, dropout=0.3)
self.head = nn.Linear(2 * hidden, classes)
def forward(self, tokens, lengths):
packed = nn.utils.rnn.pack_padded_sequence(self.emb(tokens), lengths.cpu(),
batch_first=True, enforce_sorted=False)
_, (h, _) = self.lstm(packed)
final = torch.cat([h[-2], h[-1]], dim=1) # last layer, forward + backward
return self.head(final)
model = SentimentLSTM(vocab=20000)
tokens = torch.randint(1, 20000, (4, 30)); lengths = torch.tensor([30, 22, 15, 8])
print(model(tokens, lengths).shape) # (4, 2)
print("parameters:", sum(p.numel() for p in model.parameters()))Note packed sequences: they let the LSTM skip padding tokens so the final state reflects each sequence's true end.
A memory test#
A classic diagnostic is the adding problem: given a long sequence of random numbers with two marked positions, output the sum of the two marked numbers. A vanilla RNN fails once sequences exceed a few dozen steps; an LSTM solves it for hundreds of steps โ direct evidence of long-term memory.
Variants#
- Peephole connections let gates look at the cell state.
- Coupled forget/input gates: $\mathbf{i}_t = 1 - \mathbf{f}_t$.
- GRU โ a simpler two-gate design (next lecture).
- ConvLSTM โ convolutions instead of matrix multiplications, for spatio-temporal data like weather radar.
- xLSTM (2024) revisits the LSTM with exponential gating and matrix memories to compete with transformers at scale.
Where LSTMs stand today#
Transformers replaced LSTMs for most large-scale NLP because they parallelise over sequence length and model long-range interactions directly. But LSTMs remain practical for streaming and on-device applications, small datasets, time-series forecasting and control, where their constant per-step cost and compactness are advantages.