Chuyển đến nội dung chính

レッスン 7: RNN と LSTM — 逐次シーケンス処理

リカレント ニューラル ネットワーク: アーキテクチャ、時間の経過によるバックプロパゲーション。勾配消失問題。 LSTM: セル状態、ゲート (忘れ、入力、出力)。 GRU: 簡略化されたバリアント。双方向 RNN。 PyTorch を使用した実践的なテキスト分類。

🧠 AI と ML — レッスン 6 レッスン 7: RNN と LSTM — 逐次シーケンス処理

NLP の基礎から上級まで: 自然言語処理をマスターする

パート 3: NLP のための深層学習 — RNN、LSTM、Transformer へ

xdev.asia

はじめに

リカレント ニューラル ネットワーク (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 の最大の転換点: モデルが最も重要な部分に「集中」できるようにします。