跳到主要內容

Day 35 RNN 與 LSTM 原理

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 影評資料集當作範例,示範從前處理、模型訓練到評估的完整流程。

留言

這個網誌中的熱門文章

Day 2 變數與資料型別

Day 2 變數與資料型別 引言 寫程式的過程中,變數與資料型別是處理資料的基礎。變數是存放資料的容器,資料型別則決定這筆資料有哪些特性、可以進行哪些操作。學會定義變數、認識各種資料型別,是學好 Python 的關鍵一步。 這篇文章會帶你了解 Python 中變數的觀念、如何定義變數,以及常見的資料型別,包括整數、浮點數、字串、布林值,還有串列、元組、字典與集合等容器型別。我們也會介紹變數的命名規則與撰寫風格建議,以及如何用 type() 檢查資料型別。 什麼是變數?如何在 Python 中定義變數 變數是在程式執行時用來存放資料的名稱。透過定義變數,我們可以給一筆資料一個名字,並在程式的其他地方用這個名字取用該筆資料。在 Python 中,變數不需要事先宣告型別,因為 Python 是動態型別語言,變數的型別由指定給它的值決定。 定義變數的基本語法 在 Python 中定義變數非常簡單,只要用賦值符號 = 把值指定給變數即可。例如: x = 5 # 定義變數 x,並把整數 5 賦值給它 name = "Alice" # 定義變數 name,並把字串 "Alice" 賦值給它 在這裡,x 是一個變數,被賦予整數 5;name 是另一個變數,被賦予字串 "Alice"。 變數的更新與覆寫 變數的值可以修改,也就是說,我們可以在程式的不同地方給同一個變數新的值。例如: x = 10 # x 最初被賦予 10 x = 15 # x 的值現在被更新為 15 這樣就能依照需求,在程式執行過程中靈活調整變數的值。 Python 的動態型別系統 Python 和某些靜態型別語言不同,定義變數時不需要宣告型別。賦值時,Python 會根據值自動判斷變數的型別。例如: x = 5 # x 是整數 x = 3.14 # x 變成浮點數 x = "Hi" # x 變成字串 同一個變數在程式執行過程中可以存放不同型別的值,這是 Python 的彈性之一。 常見資料型別 在 Python 中,資料型別決定我們可以對變數進行哪些操作...

Python 從入門到 PyTorch 深度學習:開啟 AI 世界的大門

Python 從入門到 PyTorch 深度學習:開啟 AI 世界的大門 隨著人工智慧(AI)與深度學習(Deep Learning)快速發展,越來越多人對這些技術產生興趣。不論你是想踏入 AI 領域的初學者,還是已經有程式基礎的開發者,學好 Python 與深度學習框架(例如 PyTorch),都能為你打開更多可能。 為什麼選擇 Python? Python 已經是資料科學與人工智慧領域的首選語言。它的語法簡潔、容易上手,而且擁有龐大的生態系與大量開源函式庫。無論是資料處理、資料視覺化,還是建立機器學習與深度學習模型,Python 都能勝任。對想進入 AI 或資料科學領域的人來說,它幾乎是必備工具。 PyTorch 是什麼? PyTorch 是由 Meta(原 Facebook)AI 研究團隊開發的開源深度學習框架,以易用、靈活和動態計算圖著稱,是許多 AI 研究人員與開發者的首選。相較於其他框架,PyTorch 的寫法更貼近原生 Python,對初學者相對友善。無論是簡單的實驗,還是複雜的深度學習模型,PyTorch 都能提供強大的支援。 這個系列能帶給你什麼? 這個系列會從 Python 的基礎開始,帶你一步一步學習,最後能自己用 PyTorch 建立深度學習模型。即使你完全沒有寫過程式,也能跟著文章的節奏累積技能,理解 AI 與深度學習的核心觀念。 本系列涵蓋的主題 Python 基礎:從變數、條件判斷到函式與模組。 資料處理工具:用 NumPy 與 Pandas 有效率地操作資料。 資料視覺化:用 Matplotlib 與 Seaborn 把資料畫成圖表。 深度學習的數學基礎:線性代數、微積分與機率。 PyTorch 入門:理解張量、模型建構與 GPU 加速。 基礎深度學習模型:CNN 與 RNN 的實作應用。 深度學習專案實戰:從資料前處理到模型部署的端到端流程。 誰適合這個系列? 程式初學者 :如果你對 AI 充滿好奇,卻還沒寫過程式,系列的第一部分會帶你快速上手 Python,並幫助你理解深度學習的基本觀念。 資料科學愛好者 :如果你已經熟悉一些資料處理方法,進階部分會教你如何用 PyTorch 建構深度學習模型。 開發者與研究人員 :想更深入了...

Day 1 Python 簡介與環境設定

Day 1 Python 簡介與環境設定 引言 在現在的科技環境裡,程式設計已經是一項重要技能。無論你是對資料科學有興趣、想成為開發者,或是想踏入人工智慧(AI)領域,學會寫程式都能明顯提升你的競爭力。在眾多程式語言中,Python 因為語法簡單、功能強大、應用範圍廣泛,成為許多人進入程式世界的第一選擇。這篇文章會帶你認識 Python 的背景與優勢,並一步步教你在不同系統上安裝與設定 Python 開發環境,最後寫出第一支 Python 程式。 為什麼選擇 Python? Python 是一種高階程式語言,由 Guido van Rossum 在 1991 年發布。Python 的設計哲學強調程式碼的可讀性,並用縮排來定義程式區塊,這點和許多使用大括號的語言不同。簡潔的語法讓它成為初學者的理想選擇;就算是經驗豐富的開發者,也能用它完成複雜的專案。 Python 的優勢如下: 簡單易學 :Python 的語法清楚、結構簡潔,初學者很快就能上手。和其他語言相比,學習曲線相對平緩,不需要先弄懂一堆複雜觀念,就能開始寫程式。 應用範圍廣泛 :從資料科學、網頁開發、人工智慧、機器學習、自動化測試到網路爬蟲,Python 都有大量開源函式庫與工具支援,而且在這些領域都扮演關鍵角色。 豐富的函式庫與框架 :Python 的函式庫生態系非常龐大。做資料分析有 NumPy、Pandas;開發網站有 Django、Flask;做深度學習有 TensorFlow、PyTorch。各種需求幾乎都能找到對應的套件,讓開發更有效率。 跨平台支援 :Python 支援 Windows、macOS、Linux 等作業系統,程式通常不需要太多修改就能跨平台執行,讓開發與部署更有彈性。 活躍的社群 :Python 擁有龐大的開發者社群。學習或開發上遇到問題,幾乎都能在社群與論壇(例如 Stack Overflow)找到答案,對初學者來說是很強的後盾,也能減少卡關時的挫折感。 Python 的應用領域 Python 的流行與強大功能,讓許多領域都開始大量使用它。以下是幾個常見的應用方向: 資料科學 :隨著大數據與人工智慧興起,資料科學大量使用 Python。NumPy、Pandas 與 Matplotlib 等工具能處理和分析龐...