Day 36 RNN 實作
引言
上一篇說明了 RNN 與 LSTM 的原理,並用 nn.LSTM 示範了輸入與輸出的形狀變化。這一篇要把 LSTM 真正接到一個任務上——IMDb 影評情緒分類。我們會做一個能判斷一段英文影評是「正面」還是「負面」的二元分類模型,並把資料前處理、詞向量(Embedding)、雙向 LSTM 與分類頭組合成一個完整的訓練流程。
這個任務同時也是 NLP(自然語言處理)領域的入門經典,涵蓋了文字資料前處理、模型設計與評估等核心步驟。完成今天的範例後,你會對「如何用 PyTorch 處理序列分類」有一個完整且可執行的樣板,未來換成中文資料或不同類別數也只需要小幅修改。對小型資料集來說,這個任務足以讓我們體驗 LSTM 的威力;但在真實大型語料上,今天這個模型架構通常會被 Transformer 取代——不過了解 LSTM 仍然是進入 NLP 世界的扎實基礎,因為它把「時間」、「記憶」、「梯度」這些概念具象化了。
準備 IMDb 資料集
PyTorch 的 torchtext 與 datasets 函式庫都內建 IMDb 資料集。為了讓範例單純一點,我們這裡用 torchtext.datasets.IMDB,它會自動下載並切成 train 與 test 兩個 split。執行第一次時需要連網下載資料,下載完成後會快取在本地,之後執行就不需要再連線。
from torchtext.datasets import IMDB
# 載入 IMDb 訓練集
train_iter = IMDB(split="train")
# 把迭代器轉成串列,方便重複取用
train_data = [(label, text) for label, text in list(train_iter)]
print("訓練樣本數:", len(train_data))
print("範例:", train_data[0][1][:120])
每筆資料是一個 (label, text) 的 tuple,label 是 1(負評)或 2(好評),text 是原始英文文字。為了符合一般習慣與損失函式預期,常會把 label 轉成 0/1,並把所有字串轉成小寫、再做基本清理(去掉 HTML 標籤、多餘空白等)。IMDb 原始資料夾帶一些 這類 HTML 殘留,建議在進詞彙表前先以正規表達式清掉,否則切詞器會把標籤當作一般字串,浪費詞彙表空間。
建立詞彙表與數值化
模型只能處理數字,所以要把文字轉成整數序列。常見做法是建立一個詞彙表(Vocabulary),把每個英文單字對應到一個唯一的整數索引。這裡用 torchtext.data.utils.get_tokenizer 切字,並用 build_vocab_from_iterator 建立詞彙表。
from torchtext.data.utils import get_tokenizer
from torchtext.vocab import build_vocab_from_iterator
tokenizer = get_tokenizer("basic_english")
def yield_tokens(data_iter):
for _, text in data_iter:
yield tokenizer(text)
# 限制詞彙表大小,未收錄字用 <unk> 表示
vocab = build_vocab_from_iterator(
yield_tokens(train_data),
max_size=20000,
specials=["<unk>", "<pad>"],
)
vocab.set_default_index(vocab["<unk>"])
print("詞彙表大小:", len(vocab))
print("'good' 對應索引:", vocab["good"])
print("'terrible' 對應索引:", vocab["terrible"])
建立詞彙表後,要把每段影評轉成等長的整數序列。這裡設定 max_len=200,超過就截斷、不足就用 <pad> 補滿,這樣模型才能用 batch 平行處理。為了讓模型有效學習,<pad> 在 LSTM 中通常會用 nn.Embedding(padding_idx=...) 處理,避免 padding 影響隱藏狀態的計算。另外,把 <unk> 與 <pad> 列在 specials 的最前面並固定索引(這裡分別是 0 與 1),後面程式讀起來會比較直覺。
順帶補充:把詞彙表大小控制在 20,000 上下,是 IMDb 這種規模任務的常見設定。詞彙表太大會讓 Embedding 參數膨脹、增加過擬合風險;太小則會讓太多字變成 <unk>,模型失去關鍵字線索。如果未來換到更專業的領域(例如法律或醫療),可以考慮搭配領域辭典或預訓練詞向量(如 GloVe、fastText),先用領域語料擴充詞彙表,再開始訓練。
建立 DataLoader
資料切分完成後,要用 DataLoader 把它包成 batch。我們順手切出 10% 當作驗證集,監控訓練過程中的過擬合情形。
import random
import torch
from torch.utils.data import DataLoader, Dataset
random.seed(42)
random.shuffle(train_data)
val_size = int(len(train_data) * 0.1)
val_data = train_data[:val_size]
split_train = train_data[val_size:]
MAX_LEN = 200
class IMDBDataset(Dataset):
def __init__(self, data, vocab, max_len=MAX_LEN):
self.data = data
self.vocab = vocab
self.max_len = max_len
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
label, text = self.data[idx]
tokens = tokenizer(text)[: self.max_len]
ids = self.vocab(tokens) + [1] * (self.max_len - len(tokens)) # 1 是 <pad> 的索引
# 把 label 1/2 轉成 0/1
return torch.tensor(ids), torch.tensor(label - 1)
train_loader = DataLoader(IMDBDataset(split_train, vocab), batch_size=64, shuffle=True)
val_loader = DataLoader(IMDBDataset(val_data, vocab), batch_size=64)
幾個小細節:把 padding 的索引固定設成 1(對應 <pad>),後面 Embedding 層會用 padding_idx=1 告訴模型把 padding 位置的梯度忽略;label 由原本的 1/2 改成 0/1,方便 nn.CrossEntropyLoss 與模型預測對齊。如果想讓訓練更穩定,可以再加入 collate_fn,統一處理不同 batch 中 padding 的長度,但對固定長度的情境來說,先把資料補到固定 200 個 token 已經能讓範例維持在好讀的大小。
用 nn.LSTM 建立情緒分類模型
模型結構很典型:Embedding → LSTM → 取最後時間點隱藏狀態 → 全連接分類頭。這裡用雙向 LSTM,並把兩個方向的隱藏狀態接在一起;最後用 nn.Linear 輸出 2 個類別。
import torch.nn as nn
class LSTMClassifier(nn.Module):
def __init__(self, vocab_size, embed_dim=128, hidden_dim=128, num_classes=2):
super().__init__()
self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=1)
self.lstm = nn.LSTM(
input_size=embed_dim,
hidden_size=hidden_dim,
num_layers=1,
bidirectional=True,
batch_first=True,
)
self.fc = nn.Linear(hidden_dim * 2, num_classes)
self.dropout = nn.Dropout(0.3)
def forward(self, x):
# x: (batch, seq_len)
emb = self.embedding(x) # (batch, seq_len, embed_dim)
out, _ = self.lstm(emb) # (batch, seq_len, hidden_dim*2)
# 取最後一個時間點的隱藏狀態
last = out[:, -1, :]
last = self.dropout(last)
return self.fc(last)
device = "cuda" if torch.cuda.is_available() else "cpu"
model = LSTMClassifier(len(vocab)).to(device)
print(model)
幾個值得注意的地方:第一,nn.Embedding 的 padding_idx=1 會讓 padding 位置的向量與梯度都不會被更新,是處理 padding 的標準做法;第二,雙向 LSTM 的隱藏維度會變成兩倍,所以 nn.Linear 的輸入要寫成 hidden_dim * 2;第三,Dropout 放在最後一個時間點的隱藏狀態之後,能進一步降低過擬合的風險。實務上也可以考慮在 Embedding 後再加一層 Dropout,或把 LSTM 換成兩層並搭配層與層之間的 dropout,這些變體都值得在後續的調參階段嘗試。
另一個常見的進階設計,是用「attention pooling」取代「只取最後一個時間點」。具體做法是給每個時間點的隱藏狀態算一個分數(attention score),加權平均後得到一個向量,這種做法對長影評特別有用,因為模型可以主動「關注」關鍵字,而不是被最後幾個 token 綁架。對 LSTM 來說,這個小技巧常常能穩定提升 1% 到 2% 的正確率。
訓練與評估
訓練迴圈的概念和 Day 30 介紹的訓練流程一樣,只是資料換成了整數序列。注意在呼叫模型前,要把 batch 也搬到對應的裝置上。
import torch.optim as optim
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=1e-3)
for epoch in range(3):
model.train()
total_loss = 0.0
for ids, labels in train_loader:
ids, labels = ids.to(device), labels.to(device)
optimizer.zero_grad()
logits = model(ids)
loss = criterion(logits, labels)
loss.backward()
optimizer.step()
total_loss += loss.item() * ids.size(0)
# 驗證
model.eval()
correct, total = 0, 0
with torch.no_grad():
for ids, labels in val_loader:
ids, labels = ids.to(device), labels.to(device)
preds = model(ids).argmax(dim=1)
correct += (preds == labels).sum().item()
total += labels.size(0)
print(f"Epoch {epoch+1}: train_loss={total_loss/len(split_train):.4f}, val_acc={correct/total:.4f}")
執行後大致會看到每個 epoch 的驗證正確率逐步上升,三個 epoch 後通常能到 85% 以上。把 epoch 數拉高、加寬隱藏維度、或換成雙層 LSTM 通常還能再往上推。不過在 CPU 上訓練這個模型會比較慢,有 GPU 的話別忘了把 device 切到 cuda。
額外提醒:LSTM 對超參數(學習率、batch 大小、hidden 維度)相當敏感,如果發現驗證正確率卡住,可以試著降低學習率、加入梯度裁剪(torch.nn.utils.clip_grad_norm_)防止梯度爆炸,或把 Embedding 維度調大。這些調參眉角會在 Day 38 進一步展開。
另一個評估上常見的疑問是:「為什麼 val_acc 不是 1?」這就要回到任務的本質——影評本身就有模糊地帶,例如「演技不錯但劇情普通」這種評論,連人類標註者也可能意見分歧。我們能追求的,是模型在大多數樣本上能抓到一致的訊號,而不是逐字精準命中每一筆。當正確率穩定在 85% 以上,搭配混淆矩陣看看主要的錯誤類型,就能判斷模型是否已經達到可用的成熟度。混淆矩陣的繪製與解讀會在 Day 39 完整介紹。
結語
這一篇把 LSTM 套到了 IMDb 情緒分類上,從資料前處理、詞彙表建立、DataLoader、模型設計到訓練迴圈走完了一輪完整流程。雖然這個範例仍偏小型,但已經具備處理真實文字分類任務的所有基本元件,下一步不管是換成中文(搭配 torchtext 的中文分詞器或自建的辭典)或是改用更深的雙層 LSTM,架構都不需要大幅更動。
明天,我們會把前面學到的所有東西整合起來,啟動一個完整的 PyTorch 專案:定義任務題目、整理資料、建立訓練與驗證流程,為接下來三天的模型設計、評估與部署鋪路。
留言
張貼留言