Day 35 RNN 與 LSTM 原理
引言
前面我們處理的都是結構固定的資料:影像可以看成二維的像素網格,表格資料可以看成列與欄的對應。這一篇要進入另一個常見的資料型態——序列資料(Sequential Data)。文字、語音、股價、心電圖這類資料,都有明顯的先後順序,前後之間也存在著重要的相依關係。
這一篇會從遞迴神經網路(Recurrent Neural Network,RNN)出發,說明它怎麼把「過去的資訊」傳遞到「現在」,接著解釋為什麼傳統 RNN 在處理長序列時會遇到梯度消失問題,並介紹為了解決這個問題而設計的長短期記憶網路(Long Short-Term Memory,LSTM)。完成今天的學習後,你就能理解序列模型背後的核心觀念,並看懂 PyTorch 中 nn.RNN 與 nn.LSTM 的設計邏輯。
為什麼需要 RNN?
回想一下,前幾篇介紹的全連接層(Fully Connected Layer)與卷積層(Convolutional Layer)都有一個共同特性:輸入與輸出彼此獨立。給定一張圖片,模型會輸出一個預測;給下一張圖片,模型不會記得上一張看過什麼。
但序列資料不一樣。一句話的意思,往往要靠上下文才讀得出來;明天的股價也和過去幾天的走勢密切相關。RNN 的設計,就是在神經網路中加入一個「記憶狀態(hidden state)」,讓模型在處理序列中的某一筆資料時,能參考前面已經看過的內容。這種「把過去帶到現在」的能力,正是 RNN 與其他網路最大的差別。
具體來說,RNN 會在每個時間點接收兩個輸入:當下的資料與上一個時間點的隱藏狀態,並輸出一個新的隱藏狀態。這個狀態會一路傳下去,成為下一個時間點的「記憶」。這樣的遞迴結構讓 RNN 理論上能處理任意長度的序列,同時在每個時間點共用同一組參數,因此參數量不會隨序列長度線性成長——這也是 RNN 在序列任務上非常有效率的原因之一。
RNN 的數學表示
把上面的概念寫成數學式會更清楚。假設輸入序列是 x₁、x₂、…、x_T,隱藏狀態是 h₁、h₂、…、h_T,那麼 RNN 的更新規則可以寫成:
# RNN 在第 t 個時間點的更新(簡化版)
h_t = tanh(W_xh @ x_t + W_hh @ h_{t-1} + b_h)
y_t = W_hy @ h_t + b_y
其中 W_xh 是輸入到隱藏狀態的權重、W_hh 是隱藏狀態之間的權重、W_hy 是隱藏狀態到輸出的權重;tanh 是雙曲正切激活函式,把數值壓在 -1 到 1 之間。h_{t-1} 是上一個時間點的隱藏狀態,這個「迴圈」就是 RNN 名字裡 Recurrent 的由來。
把這個公式沿著時間展開,就可以想像成一條很長的鏈:每一個時間點都重複同樣的運算,但帶著不同的輸入與狀態。在 PyTorch 中,我們不需要手寫這個迴圈,nn.RNN 會自動處理整條序列。
長序列的難題:梯度消失與爆炸
傳統 RNN 在訓練時,會透過反向傳播把誤差從序列尾端一路往前傳。這個過程其實就是把很多個矩陣(W_hh)反覆相乘,當序列很長的時候,這些相乘會出現兩個極端:
- 梯度消失(Vanishing Gradient):權重的數值偏小(嚴格來說是矩陣的奇異值普遍小於 1)時,反覆相乘會讓梯度快速逼近 0,導致前面的時間點幾乎學不到東西。
- 梯度爆炸(Exploding Gradient):反過來,當這些數值普遍大於 1 時,反覆相乘會讓梯度暴增,訓練參數會被大幅震盪。
梯度消失的結果是,模型「看不到」較遠的上下文。一個處理長句子的 RNN,可能讀到句尾時已經忘了開頭在講什麼。為了解決這個問題,研究者提出了 LSTM 與 GRU(Gated Recurrent Unit),它們的核心想法是用「門(gate)」來控制資訊的保留與遺忘。
另一個常見的對策是梯度裁剪(Gradient Clipping):當梯度的範數超過某個閾值時,就把它縮回合理範圍。這對梯度爆炸非常有效——實務上只要寫 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0) 就能避免訓練震盪。搭配 Adam 等自適應學習率演算法,傳統 RNN 在中等長度的序列上其實還是有不錯的表現。
LSTM:用門機制管理記憶
LSTM 是 RNN 的一個改良版本,由 Hochreiter 與 Schmidhuber 在 1997 年提出。它的關鍵設計是引入一個細胞狀態(cell state),以及三個控制資訊流的門:
- 遺忘門(Forget Gate):決定上一個細胞狀態的哪些資訊要留下、哪些要丟掉。
- 輸入門(Input Gate):決定當下的輸入有哪些新資訊要被寫入細胞狀態。
- 輸出門(Output Gate):決定細胞狀態的哪些部分要被輸出成隱藏狀態。
把這些門的運算整理成程式碼,可以更直觀地看到資料怎麼流動:
import torch
import torch.nn as nn
# 用 PyTorch 內建的 LSTM,示範如何處理序列資料
embedding = nn.Embedding(num_embeddings=10000, embedding_dim=128)
lstm = nn.LSTM(input_size=128, hidden_size=256, num_layers=1, batch_first=True)
# x 的形狀:(batch, seq_len) 的整數索引
x = torch.randint(0, 10000, (4, 30)) # 4 個樣本,每個長度 30
emb = embedding(x) # (4, 30, 128)
# 一次送入整個序列,輸出與最終狀態
output, (h_n, c_n) = lstm(emb)
print("每個時間點的輸出:", output.shape) # (4, 30, 256)
print("最終隱藏狀態 :", h_n.shape) # (1, 4, 256)
print("最終細胞狀態 :", c_n.shape) # (1, 4, 256)
從輸出可以看到,LSTM 不只回傳了每個時間點的隱藏狀態(output),還另外保留了最終的隱藏狀態(h_n)與細胞狀態(c_n)。c_n 就是 LSTM 的「長期記憶」,h_n 則是「短期工作記憶」。當要接分類器時,常用的做法是只取最後一個時間點的隱藏狀態(output[:, -1, :]),或者把整段序列做池化(pooling)。
雙向 LSTM 與多層架構
實際應用上,序列的上下文往往不只來自過去,也來自未來。舉例來說,在命名實體辨識(Named Entity Recognition)裡,要判斷一個字是不是人名的一部分,常常需要看前後的字。雙向 LSTM(Bidirectional LSTM)正是為了解決這個問題:它跑兩套 LSTM,一個從前往後讀、一個從後往前讀,最後把兩個方向的隱藏狀態接在一起。
# 雙向 LSTM:把 bidirectional 設為 True
bi_lstm = nn.LSTM(
input_size=128,
hidden_size=256,
num_layers=2, # 堆疊兩層 LSTM
bidirectional=True,
dropout=0.3, # 層與層之間加 dropout 防止過擬合
batch_first=True,
)
output, (h_n, c_n) = bi_lstm(emb)
print("雙向 LSTM 輸出:", output.shape) # (4, 30, 512)
# 因為是雙向,hidden_size 會變成兩倍
把 LSTM 堆疊多層(num_layers > 1)也是常見做法,淺層負責局部特徵、深層負責整體語意。不過要注意,LSTM 的參數量會隨 hidden_size 與 num_layers 快速膨脹,若資料量不大,建議先從單層、hidden_size 128 或 256 開始嘗試。另一個常見的取捨是:當序列很長時,純 LSTM 仍然可能不夠力,這時就會考慮使用 Transformer 家族模型(BERT、GPT 等),它們用注意力機制取代了遞迴結構,能更直接地處理長距離的相依關係。理解 LSTM 之後再去學 Transformer,會發現很多概念是相通的——例如「門」就對應到注意力中的加權機制。
結語
這一篇從 RNN 的遞迴結構講起,一路談到梯度消失、LSTM 的門機制與雙向 LSTM 的設計。我們也實際用 PyTorch 的 nn.LSTM 示範了如何一次處理整段序列,並理解 hidden state 與 cell state 的角色。掌握這些觀念後,下一篇就可以把 LSTM 套到真實的任務上——文字分類——體驗完整的訓練流程。
明天,我們會用 LSTM 實作一個簡易的文字分類模型,把 IMDb 影評資料集當作範例,示範從前處理、模型訓練到評估的完整流程。
留言
張貼留言