Day 27 自動微分 autograd
引言
Day 26 我們認識了 PyTorch 張量,知道怎麼建立、操作張量,也能在 CPU 與 GPU 之間切換裝置。不過張量本身只是「數值的容器」,還沒有辦法讓模型自己從錯誤中學習。要讓神經網路自動調整參數,必須仰賴一套機制:算出「這個參數對最後的損失影響有多大」,再沿著反方向把影響回傳給每一層。
PyTorch 把這個機制封裝在 autograd 模組裡。只要把張量標記成需要梯度,PyTorch 就會在運算的過程中建立計算圖;呼叫 backward() 之後,每個參數的梯度就會自動算出來。今天會用三個小節,把 autograd 的核心觀念一次說清楚,並用實際可執行的程式示範。
計算圖與 requires_grad
深度學習的訓練,本質上是一連串的數學運算。給定輸入 x、權重 w,經過一連串乘法、加法與激活函式,最後得到預測值;再和真實答案比較,得到損失。為了讓權重沿著「讓損失下降」的方向調整,我們需要損失對每個權重的偏導數,也就是梯度。
手算這些偏導數非常痛苦。神經網路可能有上百萬個參數,每一層的運算都不一樣。autograd 的解法是動態建立一張「計算圖」,把每一步運算記錄下來;之後只要從末端往回走,就能用鏈鎖法則把每個節點的梯度算出來。
在 PyTorch 裡,啟動 autograd 的方法很簡單:建立張量時,把 requires_grad=True 設為 True。
import torch
x = torch.tensor(3.0)
w = torch.tensor(2.0, requires_grad=True)
b = torch.tensor(1.0, requires_grad=True)
y = w * x + b # y = 2 * 3 + 1 = 7
print(y) # 輸出:tensor(7., grad_fn=<AddBackward0>)
這段程式做了三件事:建立輸入 x、權重 w 與偏差 b;w 和 b 都標記成需要梯度;接著做線性運算得到預測值。印出的結果帶有 grad_fn,表示這個張量是由某個運算產生的,而且這個運算的「來源」已被記錄下來,這就是計算圖的起點。
另外要注意:requires_grad 一旦設成 True,只要這個張量參與運算,結果(也叫「子張量」)就會自動繼承這個屬性。
print(y.requires_grad) # 輸出:True
print(y.grad_fn) # 輸出:<AddBackward0 object at 0x...>
可以用 grad_fn 屬性觀察計算圖的結構;如果一個張量是直接建立的(例如 torch.tensor(2.0)),這個屬性會是 None。
反向傳播:呼叫 backward() 取得梯度
計算圖建好之後,要怎麼把梯度算出來?答案是呼叫末端張量的 backward() 方法。PyTorch 會從這個張量開始,沿著計算圖往回走,把每一個需要梯度的張量的梯度填到 .grad 屬性。
先做一個最簡單的範例:假設損失是 y 本身,也就是 L = y = w * x + b。對 w 和 b 取偏導數,預期結果分別是 x(3)與 1。
import torch
x = torch.tensor(3.0)
w = torch.tensor(2.0, requires_grad=True)
b = torch.tensor(1.0, requires_grad=True)
y = w * x + b # 損失 L = y
y.backward() # 反向傳播
print(w.grad) # 輸出:tensor(3.)
print(b.grad) # 輸出:tensor(1.)
print(x.grad) # 這裡 x 沒有 requires_grad,所以是 None
w 的梯度是 3(因為 L 對 w 的偏導數是 x),b 的梯度是 1,這個結果和手算一致。x 本身沒有標記成需要梯度,所以不會有 grad 屬性,這也是正確的行為,因為在訓練時只有模型參數需要更新。
實務上損失很少只是一個純量,它通常是許多筆資料的誤差平均。PyTorch 的 backward() 預期收到的引數是一個純量張量;如果末端張量不是純量,就要傳入一個 gradient 引數,告訴 autograd 怎麼把多維結果「加權」成純量再反向傳播。常見的做法是先用損失函式把輸出收斂成純量,這部分會在 Day 30 詳細說明。
另外,backward() 預設會把新的梯度「累加」到現有的 .grad 上。也就是說,第二次呼叫 backward() 之前,必須先呼叫 optimizer.zero_grad() 或 w.grad.zero_() 把梯度清空,否則梯度會越加越多。
import torch
w = torch.tensor(2.0, requires_grad=True)
x = torch.tensor(3.0)
y1 = w * x
y1.backward()
print(w.grad) # 輸出:tensor(3.)
w.grad.zero_() # 把梯度清零
y2 = w * x
y2.backward()
print(w.grad) # 輸出:tensor(3.)
停止追蹤梯度:no_grad 與 detach
並非所有運算都需要建立計算圖。在訓練時我們想保留梯度以便更新參數,但在「驗證」或「推論」階段,只需要模型輸出的數值;如果這時還建立計算圖,不但浪費記憶體,也可能造成額外的麻煩。PyTorch 提供兩個常用的方法:torch.no_grad() 與 tensor.detach()。
torch.no_grad() 是一個情境管理器(context manager),在它包住的範圍內,所有運算都不會被 autograd 記錄。
import torch
w = torch.tensor(2.0, requires_grad=True)
with torch.no_grad():
y = w * 3 # y 不會建立計算圖
print(y.requires_grad) # 輸出:False
print(y.grad_fn) # 輸出:None
驗證階段通常會這樣寫:model.eval() 把模型切到評估模式(會關閉 dropout 之類的隨機行為),再用 torch.no_grad() 包住推論程式碼,既省記憶體也避免不小心把驗證資料寫進計算圖。
detach() 則是張量本身的方法,呼叫後會產生一個「脫離計算圖」的副本,原本張量的數值不變,但不再被 autograd 追蹤。這個方法常用於把模型輸出轉成 NumPy、或是在某些自訂的損失函式裡切斷不需要的梯度。
import torch
w = torch.tensor(2.0, requires_grad=True)
y = w * 3
y_detached = y.detach()
print(y_detached.requires_grad) # 輸出:False
print(y_detached.grad_fn) # 輸出:None
print(y_detached) # 輸出:tensor(6.)
# y_detached 已經脫離計算圖,對它呼叫 backward() 會直接報錯:
# RuntimeError: element 0 of tensors does not require grad and does not have a grad_fn
多變數案例與梯度累積
真實的損失很少只是一個純量,而是許多筆誤差的總和或平均。當末端張量是多維時,backward() 需要一個對應形狀的 gradient 引數,告訴 autograd 怎麼把這些值「加權」再反向傳播。最常見的做法是用 loss.sum() 或 loss.mean() 把結果收斂成純量,再呼叫 backward(),這樣不必額外提供 gradient。
import torch
w = torch.tensor([[1.0, 2.0]], requires_grad=True) # shape [1, 2]
x = torch.tensor([[3.0, 4.0],
[5.0, 6.0]]) # shape [2, 2]
# 對每一筆樣本算一個損失,得到 [2] 的一維張量
loss = ((w * x).sum(dim=1) - torch.tensor([5.0, 15.0])) ** 2
# 末端是多維,用 mean() 收斂成純量後再 backward()
loss.mean().backward()
print(w.grad) # 輸出:tensor([[28., 36.]])
另一個重點是,backward() 預設會把新的梯度累加到 .grad 上,而不是覆寫。這在某些進階情境(例如梯度累積模擬大批次)很實用,但在大多數訓練流程裡會造成「梯度越來越大的 bug」,務必記得在每一步訓練前呼叫 optimizer.zero_grad() 來清空。
結語
今天我們理解了 autograd 的核心:把運算過程建成計算圖,再用反向傳播把每個參數的梯度算出來。requires_grad 決定哪些張量要被記進計算圖,backward() 從末端往回算梯度,no_grad() 與 detach() 則提供停止追蹤的機制。這些觀念是之後所有訓練流程的基礎,後面寫到的損失函式、最佳化器、訓練迴圈,背後都依賴 autograd。順帶記住兩件事:backward() 預期收到的是純量張量、每次呼叫前要記得清空梯度,這兩點在 Day 30 寫訓練迴圈時會再次出現。
明天,我們會把資料這一塊補齊。實際訓練模型時,往往需要從磁碟讀取大量資料、做預處理、切成小批次。PyTorch 提供 torch.utils.data.Dataset 與 DataLoader 把這些工作標準化,是進入 MNIST 實作前不可缺的準備。
留言
張貼留言