π Long Short-Term Memory (LSTM)
Descriptionβ
< What is it? >β
A Long Short-Term Memory (LSTM) network is a gated type of RNN. It was designed to preserve useful information and gradient flow across longer sequences than a vanilla RNN can usually handle.
An LSTM cell carries two vectors from one time step to the next:
- : the cell state, a long-running memory path
- : the hidden state, the output exposed to the next layer or time step
Key pointsβ
< Gates control memory >β
At each time step, an LSTM uses learned sigmoid gates to decide what to forget, write, and expose. Let
The gates and state updates are:
Here, gives values between and , and means element-wise multiplication.
| Gate | Simple interpretation |
|---|---|
| Forget gate | How much of the previous memory to keep |
| Input gate | How much new candidate information to write |
| Output gate | How much of the cell memory to reveal as |
< Why LSTM helps with long sequences >β
The cell-state update has an additive path:
When the forget gate is near , useful memory and its gradient can pass through many time steps with less shrinkage than in a vanilla RNN. Gates can also discard irrelevant information when it is no longer useful.
LSTM reduces the vanishing-gradient problem; it does not make every long dependency easy or guarantee that gradients never explode.
- Example: after reading βThe animal rested because it was tired,β an LSTM can choose to retain information about
animaluntil it processesit.
< Inputs and outputs >β
For a batch-first PyTorch input with shape (B, T, D):
| Symbol | Meaning | Typical shape |
|---|---|---|
| Batch size | β | |
| Number of time steps | β | |
| Input feature or embedding size | β | |
| LSTM hidden size | β | |
| Output | Hidden state for every step | (B, T, H) |
| Final states | Final hidden and cell states | (L, B, H) each |
Here, is the number of stacked LSTM layers for a one-directional LSTM.
< In PyTorch >β
import torch
import torch.nn as nn
# Input shape: (batch, time_steps, embedding_size)
x = torch.randn(32, 20, 128)
lstm = nn.LSTM(
input_size=128,
hidden_size=256,
batch_first=True,
)
output, (h_n, c_n) = lstm(x)
# output: (32, 20, 256), one hidden state per time step
# h_n: (1, 32, 256), final hidden state
# c_n: (1, 32, 256), final cell state
< Where it is used >β
LSTMs remain useful for modest-size time-series, speech, sensor, and sequence tasks, especially when sequential processing is acceptable. For very long text or large-scale language modeling, Transformers usually provide better parallelism and global-context modeling.
Comparisonβ
< LSTM vs GRU >β
| LSTM | GRU | |
|---|---|---|
| State | Separate hidden state and cell state | One hidden state |
| Gates | Forget, input, output | Update and reset |
| Parameters | More | Fewer |
| Practical choice | Useful when explicit cell memory is helpful | Often a simpler, faster gated baseline |
Both are gated RNNs and are usually more capable than a vanilla RNN for long dependencies.
Related ideasβ
- Recurrent Neural Network (RNN) introduces recurrence and hidden states.
- Embeddings explains how text tokens become LSTM inputs.
- Vanishing & Exploding Gradients explains why gated cells are helpful.
- Transformer provides a non-recurrent alternative for sequence modeling.