はじめに
リカレント ニューラル ネットワーク (RNN) は、シーケンシャル データ (テキスト、時系列、オーディオ) 用に設計された初のニューラル ネットワーク アーキテクチャです。 Transformer は多くの問題で置き換えられてきましたが、RNN/LSTM を理解することは、tention と 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 PyTorch を使用した LSTM
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倍の計算 |
📌 なぜTransformerが必要なのでしょうか? RNN/LSTMは各トークンを順番に処理します → 並列できない → 長いチェーンでは遅いです。トランスはこの問題を解決します。
次の記事
レッスン 8: 注意のメカニズム — NLP の最大の転換点: モデルが最も重要な部分に「集中」できるようにします。