簡介
循環神經網路 (RNN) 是第一個專為 序列 資料(文字、時間序列、音訊)而設計的神經網路架構。雖然Transformer在許多問題上已經被取代,但理解RNN/LSTM是理解Attention和Transformer為何誕生的基礎。
1. RNN——循環神經網絡
1.1 架構
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 梯度消失問題
當鏈很長時,梯度隨著每個時間步呈指數下降**→RNN“忘記”鏈開頭的信息。
"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——長短期記憶
2.1 LSTM 單元架構
Cell State (C)
──────────────────────────────────────────
│ │ │
┌────┴────┐ ┌────┴────┐ ┌────┴────┐
│ Forget │ │ Input │ │ Output │
│ Gate │ │ Gate │ │ Gate │
│ σ(Wf) │ │ σ(Wi) │ │ σ(Wo) │
└─────────┘ └─────────┘ └─────────┘
| 門 | 食譜 | 功能 |
|---|---|---|
| 忘記 | $f_t = \sigma(W_f \cdot [h_{t-1}, x_t] + b_f)$ | 忘了什麼? |
| 輸入 | $i_t = \sigma(W_i \cdot [h_{t-1}, x_t] + b_i)$ | 什麼是新的? |
| 輸出 | $o_t = \sigma(W_o \cdot [h_{t-1}, x_t] + b_o)$ | 什麼輸出? |
2.2 LSTM 與 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——門控循環單元
GRU 簡化了 LSTM:將遺忘門+輸入門組合成更新門,刪除單獨的狀態單元。
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)
| 比較 | LSTM | 格魯烏 |
|---|---|---|
| 參數 | 更多(4 個門) | 少(2 個門) |
| 培訓 | 慢一點 | 更快 |
| 長序列 | 更好 | 好 |
| 品質 | 通常相當於 | 通常相當於 |
4. 雙向 RNN
Forward: h₁ → h₂ → h₃ → h₄
→ concat → output
Backward: h₄ ← h₃ ← h₂ ← h₁
從兩個方向閱讀文本-理解前後的上下文。
總結
| 建築 | 優勢 | 限制 |
|---|---|---|
| 普通 RNN | 簡單 | 消失的梯度,很快就被遺忘 |
| LSTM | 遠端依賴 | 緩慢、順序(非平行) |
| 格魯烏 | 比 LSTM 更輕、更有效 | 類似 LSTM |
| 雙向 | 理解 2 向上下文 | 2x 計算 |
📌 **為什麼我們需要 Transformer? ** RNN/LSTM 順序處理每個 token → 不能並行 → 長鏈速度慢。變壓器解決了這個問題。
下一篇文章
第8課:注意力機制——NLP最大的轉折點:讓模型「聚焦」在最重要的部分。