Introduction
Recurrent Neural Networks (RNN) is the first neural network architecture designed for sequential data — text, time series, audio. Although Transformer has been replaced in many problems, understanding RNN/LSTM is the foundation to understand why Attention and Transformer were born.
1. RNN — Recurrent Neural Network
1.1 Architecture
x₁ x₂ x₃ x₄
│ │ │ │
▼ ▼ ▼ ▼
┌───────┐ ┌───────┐ ┌───────┐ ┌───────┐
h₀──│ RNN │──▶│ RNN │──▶│ RNN │──▶│ RNN │──▶ h₄
│ Cell │ │ Cell │ │ Cell │ │ Cell │
└───┬───┘ └───┬───┘ └───┬───┘ └───┬───┘
│ │ │ │
▼ ▼ ▼ ▼
y₁ y₂ y₃ y₄
$$h_t = \tanh(W_{hh} \cdot h_{t-1} + W_{xh} \cdot x_t + b)$$
1.2 Vanishing Gradient problem
When the chain is long, the gradient decreases exponentially with each timestep → RNN "forgets" the information at the beginning of the chain.
"The cat, which sat on the mat and watched the birds for hours, was ___"
↑ ↑
Thông tin cần ở rất xa Cần predict ở đây
→ Gradient ≈ 0 khi backpropagate ngược lại!
2. LSTM — Long Short-Term Memory
2.1 LSTM Cell Architecture
Cell State (C)
──────────────────────────────────────────
│ │ │
┌────┴────┐ ┌────┴────┐ ┌────┴────┐
│ Forget │ │ Input │ │ Output │
│ Gate │ │ Gate │ │ Gate │
│ σ(Wf) │ │ σ(Wi) │ │ σ(Wo) │
└─────────┘ └─────────┘ └─────────┘
| Gate | Recipe | Function |
|---|---|---|
| Forget | $f_t = \sigma(W_f \cdot [h_{t-1}, x_t] + b_f)$ | Forgot what? |
| Input | $i_t = \sigma(W_i \cdot [h_{t-1}, x_t] + b_i)$ | What's new? |
| Output | $o_t = \sigma(W_o \cdot [h_{t-1}, x_t] + b_o)$ | What output? |
2.2 LSTM with PyTorch
import torch
import torch.nn as nn
class TextClassifierLSTM(nn.Module):
def __init__(self, vocab_size, embed_dim, hidden_dim, num_classes):
super().__init__()
self.embedding = nn.Embedding(vocab_size, embed_dim)
self.lstm = nn.LSTM(
embed_dim, hidden_dim,
num_layers=2,
bidirectional=True,
batch_first=True,
dropout=0.3,
)
self.fc = nn.Linear(hidden_dim * 2, num_classes) # *2 for bidirectional
self.dropout = nn.Dropout(0.3)
def forward(self, x):
# x: (batch_size, seq_len)
embedded = self.embedding(x) # (B, L, embed_dim)
output, (hidden, cell) = self.lstm(embedded) # output: (B, L, hidden*2)
# Lấy hidden state cuối cùng từ 2 directions
hidden = torch.cat((hidden[-2], hidden[-1]), dim=1) # (B, hidden*2)
hidden = self.dropout(hidden)
return self.fc(hidden) # (B, num_classes)
# Khởi tạo
model = TextClassifierLSTM(
vocab_size=30000,
embed_dim=128,
hidden_dim=256,
num_classes=3,
)
3. GRU — Gated Recurrent Unit
GRU simplifies LSTM: combine forget + input gate into update gate, removing separate state cell.
class TextClassifierGRU(nn.Module):
def __init__(self, vocab_size, embed_dim, hidden_dim, num_classes):
super().__init__()
self.embedding = nn.Embedding(vocab_size, embed_dim)
self.gru = nn.GRU(
embed_dim, hidden_dim,
num_layers=2,
bidirectional=True,
batch_first=True,
dropout=0.3,
)
self.fc = nn.Linear(hidden_dim * 2, num_classes)
def forward(self, x):
embedded = self.embedding(x)
output, hidden = self.gru(embedded) # Không có cell state
hidden = torch.cat((hidden[-2], hidden[-1]), dim=1)
return self.fc(hidden)
| Compare | LSTM | GRU |
|---|---|---|
| Parameters | More (4 gates) | Less (2 gates) |
| Training | Slower | Faster |
| Long sequences | Better | Good |
| Quality | Usually equivalent to | Usually equivalent to |
4. Bidirectional RNN
Forward: h₁ → h₂ → h₃ → h₄
→ concat → output
Backward: h₄ ← h₃ ← h₂ ← h₁
Read text from both directions — understand the context before AND after.
Summary
| Architecture | Advantages | Limitations |
|---|---|---|
| Vanilla RNN | Simple | Vanishing gradient, quickly forgotten |
| LSTM | Long-range dependencies | Slow, sequential (not parallel) |
| GRU | Lighter than LSTM, effective | Similar to LSTM |
| Bidirectional | Understanding 2-way context | 2x computation |
📌 Why do we need a Transformer? RNN/LSTM processes each token sequentially → cannot be parallel → slow with long chains. Transformer solves this problem.
Next article
Lesson 8: Attention Mechanism — The biggest turning point of NLP: allowing the model to "focus" on the most important part.