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

Lesson 7: RNN & LSTM — Sequential Sequence Processing

Recurrent Neural Networks: architecture, backpropagation through time. Vanishing gradient problem. LSTM: cell state, gates (forget, input, output). GRU: simplified variant. Bidirectional RNN. Hands-on text classification with PyTorch.

🧠 AI & ML — Lesson 6 Lesson 7: RNN & LSTM — Sequential Sequence Processing

NLP from Basics to Advanced: Mastering Natural Language Processing

Part 3: Deep Learning for NLP — RNN, LSTM, to Transformer

xdev.asia

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)   │
    └─────────┘    └─────────┘    └─────────┘
GateRecipeFunction
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)
CompareLSTMGRU
ParametersMore (4 gates)Less (2 gates)
TrainingSlowerFaster
Long sequencesBetterGood
QualityUsually equivalent toUsually 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

ArchitectureAdvantagesLimitations
Vanilla RNNSimpleVanishing gradient, quickly forgotten
LSTMLong-range dependenciesSlow, sequential (not parallel)
GRULighter than LSTM, effectiveSimilar to LSTM
BidirectionalUnderstanding 2-way context2x 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.