Seventeen years after the LSTM, Kyunghyun Cho and colleagues (2014) introduced a streamlined alternative while developing neural machine translation: the Gated Recurrent Unit (GRU). It merges the cell and hidden states and uses two gates instead of three. It trains faster, has fewer parameters, and in many tasks matches the LSTM's accuracy.
The equations#
(Some references and libraries swap the roles of $\mathbf{z}_t$ and $1 - \mathbf{z}_t$; the meaning is the same.)
Interpreting the gates#
- Update gate $\mathbf{z}_t$ decides how much of the state to replace. When $\mathbf{z}_t \approx 0$, the state is copied forward unchanged — preserving memory and gradients, like the LSTM's forget gate set to 1. It effectively couples the LSTM's forget and input gates: whatever is written replaces an equal fraction of what is kept.
- Reset gate $\mathbf{r}_t$ decides how much of the previous state to use when computing the candidate. When $\mathbf{r}_t \approx 0$, the unit ignores the past and behaves like a feedforward layer on the current input — useful at boundaries, such as the start of a new phrase.
The gradient along the direct path is $\partial\mathbf{h}_t/\partial\mathbf{h}_{t-1} \supseteq \text{diag}(1 - \mathbf{z}_t)$, providing the same highway for long-range gradient flow as the LSTM's cell.
GRU vs LSTM#
| LSTM | GRU | |
|---|---|---|
| States | Cell $\mathbf{c}_t$ + hidden $\mathbf{h}_t$ | Hidden $\mathbf{h}_t$ only |
| Gates | Forget, input, output | Update, reset |
| Parameters (per layer) | $4(H(H + D) + H)$ | $3(H(H + D) + H)$ |
| Speed | Slower | ~25% fewer operations |
| Output gating | Yes — can hide memory | No — full state exposed |
Empirical comparisons (Chung et al., 2014; Jozefowicz et al., 2015; Greff et al., 2017) found no consistent winner: GRUs often match LSTMs, sometimes win on smaller datasets, while LSTMs can be slightly stronger on tasks requiring counting or very precise memory — the output gate and separate cell give extra control. A large study of LSTM variants concluded that the forget gate and output activation are the most critical components.
Comparing them in code#
import torch
import torch.nn as nn
def make_copy_task(batch, T, n_symbols=8, delay=50):
"""Remember a short sequence of symbols, output it after a long delay."""
seq = torch.randint(1, n_symbols, (batch, 5))
x = torch.zeros(batch, T, dtype=torch.long); x[:, :5] = seq; x[:, 5 + delay] = n_symbols # "go" marker
y = torch.zeros(batch, T, dtype=torch.long); y[:, 6 + delay:11 + delay] = seq
return x, y
class Seq(nn.Module):
def __init__(self, cell, H=64, V=10):
super().__init__()
self.emb = nn.Embedding(V, 16)
self.rnn = {"lstm": nn.LSTM, "gru": nn.GRU, "rnn": nn.RNN}[cell](16, H, batch_first=True)
self.out = nn.Linear(H, V)
def forward(self, x):
return self.out(self.rnn(self.emb(x))[0])
T = 70
for cell in ["rnn", "gru", "lstm"]:
torch.manual_seed(0); m = Seq(cell); opt = torch.optim.Adam(m.parameters(), 3e-3)
for step in range(1500):
x, y = make_copy_task(64, T)
logits = m(x)
loss = nn.functional.cross_entropy(logits[:, 56:61].reshape(-1, 10), y[:, 56:61].reshape(-1))
opt.zero_grad(); loss.backward(); nn.utils.clip_grad_norm_(m.parameters(), 1.0); opt.step()
x, y = make_copy_task(512, T)
acc = (m(x)[:, 56:61].argmax(-1) == y[:, 56:61]).float().mean().item()
print(f"{cell:<5} copy accuracy after a 50-step delay: {acc:.2f}")On such long-delay memory tasks, the gated cells typically succeed where the vanilla RNN stays near chance.
Practical guidance for recurrent models#
- Start with a GRU for speed; try an LSTM if accuracy matters and compute allows.
- Clip gradients (norm 1–5) — always.
- Use 2–3 layers with dropout between them; beyond that, gains diminish.
- Use bidirectional layers when the whole sequence is available at prediction time.
- Pack padded sequences and sort or bucket by length for efficiency.
- Consider 1-D CNNs or temporal convolutional networks — they parallelise well and often compete with RNNs on sequence classification.
- For long-context language tasks, prefer transformers or modern state-space models.