JOVANA
Explore Library Glossary Getting Started Three Levels Fields How it works Mission
Join the mission
All guides

讓梯度保持健康:正規化與穩定訓練

為什麼深層網路會學不動?看正規化、殘差與裁剪如何讓梯度保持鮮活。

網路越深,越難訓練

在前兩篇指南中,你看到了視覺模型是怎麼學習的:前向傳播做出預測,損失函數衡量它錯得多離譜,反向傳播計算梯度,最佳化器再把權重往下坡方向推一點。感覺整個食譜好像已經完成了——所以想讓網路更聰明,就多疊幾層不就好了?更多層代表更強的能力,能先認出邊緣、再認出紋理、最後認出完整的物體。理論上,越深應該總是越好。

現實卻冷酷地不同。多年來研究者發現,天真地疊層反而讓網路訓練得更差而不是更好——一個 30 層的樸素網路,最後的訓練誤差竟然可能比 15 層的還高,即使更深的那個嚴格來說能力更強。罪魁禍首在於梯度回傳途中發生的事。反向傳播把誤差訊號往回送過每一層,而在每一層,這個訊號都會被乘上一個局部因子。把許多數字連乘起來,往往會發生兩種壞事之一:乘積塌縮趨近於零,或者爆炸趨近於無限大。這就是著名的梯度消失(與梯度爆炸)問題

在動用任何數學之前,先抓住直覺。想像有一排 100 個人在玩傳話遊戲。第一個人小聲說出一句話,每個人再傳給下一個。等傳到第 100 個人時,這句話要嘛微弱到沒人聽得見(消失),要嘛被扭曲放大成一堆胡言亂語(爆炸)。梯度就正是那句悄悄話,從最後一層一路往回傳到第一層。最前面的那幾層——本該學會基本邊緣的層——收到的訊息亂到幾乎無法更新,於是它們永遠學不到任何有用的東西。

訊號在前向傳播時必須依序穿過每一層——而梯度在回傳時也必須穿過每一層。鏈條越長,訊息淡掉或爆掉的機會就越多。

卷積網路被畫成一疊層,從左邊的輸入影像,經過數個處理階段,到右邊的輸出預測。

好消息是:這是一個已被工程解決的問題。如今擁有數百層的網路能可靠地訓練,正是因為我們打造了一套工具,讓梯度在漫長旅途中保持健康。本篇指南會一塊一塊地組裝這套工具:正規化(批次正規化、層正規化)讓激活值維持在合理範圍;殘差連接給梯度一條直通回去的快速車道;而梯度裁剪則是一條安全帶,把爆炸的部分上限封住。讀完後,你不只會知道每個工具做什麼,還會精確知道它修好了「傳話遊戲」問題的哪一部分。

精確理解梯度消失與梯度爆炸

讓我們把傳話遊戲的比喻變成精確的東西。反向傳播其實就是微積分裡的連鎖律一遍又一遍地套用。想知道「如果你輕輕擾動某個前層的激活值,損失會如何改變」,你就得把這個前層到損失之間每一層的局部因子——也就是「輸出隨輸入改變多少」——全部連乘起來。因此,前層的梯度是一長串逐層因子的乘積——而許多數字相乘的行為,和相加非常不一樣。

\frac{\partial \mathcal{L}}{\partial a_{1}} \;=\; \frac{\partial \mathcal{L}}{\partial a_{n}} \;\prod_{k=2}^{n} \frac{\partial a_{k}}{\partial a_{k-1}}

第一層的梯度,等於最後一層的梯度乘上一連串逐層因子。

我們逐個符號讀。a_kk 層的激活值向量——也就是那一層輸出的數字。\mathcal{L} 是損失。左邊的 \partial \mathcal{L}/\partial a_1第一層的梯度——如果我們擾動最前面的激活值,損失會改變多少,這正是告訴第 1 層該如何更新的東西。右邊,\partial \mathcal{L}/\partial a_n最後一層的梯度(反向傳播從這裡開始),而那個大大的 \prod_{k=2}^{n} 代表「對從 2 到 n 的每一層 k,把它們全部相乘」。每個因子 \partial a_k / \partial a_{k-1} 才是核心:k 層的輸入改變時,它的輸出改變多少。如果這些因子每個都略小於 1,許多個相乘就會趨近 0——梯度消失;如果每個都略大於 1,乘積就會無限增長——梯度爆炸

現在來個讓人切身有感的算式。那些逐層因子,部分取決於激活函數的斜率。經典的 sigmoid 和 tanh 函數會飽和:當輸入很大(正或負)時,曲線變平,斜率就極小。sigmoid 的導數從不超過 0.25,而且通常還小得多。慷慨地假設,每一層剛好給你 0.25,那麼 10 層連乘就是:0.25^{10} \approx 0.0000009。抵達第 1 層的梯度,大約只剩從頂端出發時的一百萬分之一——實務上等於消失了。若疊上 20 層這樣的層,你就掉到 10^{-13}。這正是為什麼深層的 sigmoid/tanh 網路出了名地拒絕學習。

這個領域發現的第一個修法,其實就只是換個更好的激活函數。ReLU 在輸入為正時原封不動地輸出,否則輸出零,所以對任何正輸入,它的斜率剛好是 1,而不是 0.25。因子為 1 完全不會縮小乘積——把 1 自乘一百次,你還是得到 1。ReLU 並沒有完全解決所有問題(它可能讓神經元關閉,也擋不住爆炸),但它阻止了 sigmoid 造成的那種無情的幾何式衰減。這正是為什麼幾乎每個現代 CNN 都使用 ReLU 或它的近親。

ReLU 對負輸入是平的(斜率 0),但對正輸入是漂亮的 45 度直線(斜率 1)。這個為 1 的斜率,正是讓連鎖律乘積不致塌縮的關鍵因子。

ReLU 激活函數的圖:在負 x 處是貼著零的水平線,在正 x 處則是一條 45 度向上的直線。

批次正規化:在資料流動時標準化激活值

連鎖律乘積會失控,部分原因在於激活值穿過層時會漂移到不健康的範圍——有些層輸出巨大的數字,有些則輸出極小的,而且當訓練更新前面的權重時,整個分布還會不斷移動。批次正規化(BatchNorm)就是那個主力修法。想法很簡單:在網路內部選定的某一點,取每個特徵的預激活值,用當前小批次的統計量把它們強制變成零均值、單位變異數,然後再讓網路把它們重新縮放與平移。讓激活值維持在合理且一致的範圍,能讓損失地形更平滑,讓你能安全地使用較高的學習率,加速收斂,甚至附帶一點正規化的雜訊當紅利。

\hat{x}_{i} \;=\; \frac{x_{i} - \mu_{B}}{\sqrt{\sigma_{B}^{2} + \epsilon}} \,, \qquad y_{i} \;=\; \gamma\,\hat{x}_{i} + \beta

第一步:標準化為零均值、單位變異數。第二步:用兩個可學習參數重新縮放與平移。

逐個符號看:x_i 是一個預激活值(批次中某個樣本的某個特徵)。\mu_B\sigma_B^2 是針對那個特徵、在小批次 B 上計算出的均值與變異數——那個下標 B 正是重點,代表「用當前批次的統計量」。減去 \mu_B 把數值置中到零;除以 \sqrt{\sigma_B^2} 把它們縮放到單位散布。\epsilon 是加在根號底下的一個極小常數(比如 10^{-5}),確保當某個特徵碰巧是常數時也絕不會除以零。結果 \hat{x}_i 就是標準化後的值。接著是巧妙的轉折:\gamma(縮放)與 \beta(平移)是可學習的參數,最佳化器會像訓練其他權重一樣訓練它們,產生最終輸出 y_i

為什麼要讓 \gamma\beta 可學習?因為強迫每一層都剛好是零均值、單位變異數,可能會丟掉有用的資訊——有時最好的表示確實就是被平移或拉伸過的。透過給網路自己的旋鈕,BatchNorm 並不強迫一個固定的分布;它只是給網路一個健康的起始分布,網路若想要還能再調整。極端來說,如果 \gamma = \sqrt{\sigma_B^2}\beta = \mu_B,網路就能完美地把正規化還原回去。所以 BatchNorm 永遠不會損害表達能力——它只會幫助最佳化。這個「可以自由還原」的特性,正是為什麼它幾乎能安全地塞進任何地方。

import numpy as np

# One feature's pre-activations across a mini-batch of 4 examples
x = np.array([2.0, 4.0, 6.0, 8.0])

mu  = x.mean()          # 5.0  -> the batch mean (mu_B)
var = x.var()           # 5.0  -> the batch variance (sigma_B^2)
eps = 1e-5

x_hat = (x - mu) / np.sqrt(var + eps)
# x_hat ~= [-1.342, -0.447, 0.447, 1.342]  (now mean 0, variance 1)

# Learnable scale and shift, initialised to the identity (gamma=1, beta=0)
gamma, beta = 1.0, 0.0
y = gamma * x_hat + beta
# Training is free to learn other gamma/beta if a different spread helps.
一個 4 數的批次 {2,4,6,8}:均值 5、變異數 5,所以標準化後的值約為 {-1.34, -0.45, 0.45, 1.34}。
BatchNorm 把每個特徵跨小批次中的各個樣本做標準化,再套上可學習的縮放 gamma 與平移 beta。

示意圖顯示一行預激活值被置中為零均值、單位變異數,再由 gamma 與 beta 重新縮放與平移。

當批次不可靠:層正規化與它的夥伴

BatchNorm 有個藏在明處的致命弱點:它依賴批次統計量。如果你的批次大小很小——比如 2 或 4 張影像,這在每張影像很大或記憶體吃緊時很常見——那麼 \mu_B\sigma_B^2 就只能從寥寥幾個數字估計出來,變得雜訊大又不可靠。這時正規化注入的反而是隨機性而非穩定性。BatchNorm 對於變長或序列式的輸入(像文字,或一連串影格)也很彆扭,因為「批次」並不是一個乾淨固定、可供平均的形狀。

層正規化(LayerNorm)只用一次換軸就繞過了整個問題。它不再把每個特徵跨整批樣本做正規化,而是把單一樣本內所有特徵之間做正規化。每個樣本只用它自己的數字來標準化——所以批次裡有 1 個樣本還是 1000 個都無所謂,其他樣本長什麼樣也無所謂。LayerNorm 完全是與批次無關的。

\mu \;=\; \frac{1}{H}\sum_{j=1}^{H} x_{j} \,, \qquad \sigma^{2} \;=\; \frac{1}{H}\sum_{j=1}^{H}\bigl(x_{j} - \mu\bigr)^{2}

LayerNorm 的均值與變異數,是針對單一樣本的所有特徵取的。

逐個符號看:H那一個樣本的特徵數(例如該層特徵向量的長度)。總和 \sum_{j=1}^{H} 是跑過單一樣本的各個特徵 j——而不是跑過整批。所以 \mu 是那一個樣本在自己各特徵上的平均,而 \sigma^2 是這些相同特徵環繞 \mu 的散布。算完之後,你套用和 BatchNorm 完全一樣的兩個步驟:標準化 \hat{x} = (x - \mu)/\sqrt{\sigma^2 + \epsilon},再縮放與平移 y = \gamma\,\hat{x} + \beta,其中 \gamma\beta 可學習。公式是同一套正規化;只有被平均的軸換了——而正是這個改變移除了對批次的依賴。

用一張圖把這個對比記在腦中:把資料排成一個格子,列是樣本,欄是特徵。BatchNorm 沿著一欄做正規化(一個特徵,跨所有樣本)。LayerNorm 沿著一列做正規化(一個樣本,跨它所有特徵)。同一個格子,垂直的兩個方向。在這兩個極端之間,住著 GroupNormInstanceNorm,它們在單一樣本內把通道分成選定的子群來做正規化——當你想要與批次無關、又想比完整的 LayerNorm 多一點結構時,是好用的折衷。

殘差連接:梯度的高速公路

正規化讓激活值維持健康,但它並沒有完全治好深度的問題。即使到處都放了 BatchNorm,非常深的樸素網路仍會退化——過了某個臨界點,加層反而讓訓練誤差和測試誤差雙雙上升。最終讓 100 層以上的網路得以訓練的突破,簡單到幾乎令人尷尬:殘差連接(或稱跳接),這是 ResNet 的核心想法。它直接攻擊梯度消失問題——靠的是改變網路的形狀,而不只是改變激活值的尺度。

改變是這樣的。普通的一段層接收輸入 x,計算某個變換 y = F(x)。殘差區塊則改算 y = F(x) + x——它做同樣的工作 F(x),然後把原始輸入加回來。把 F(x) 想成一條塞車的市區道路,把 +x 想成一條繞過所有車流的快速車道:資訊(與梯度)永遠可以走那條直達路線繞過這個區塊,而不必在每一層裡龜速爬行。這個區塊只需要學習殘差——在「把輸入原樣傳下去」之上的那個小修正 F(x)——而這也是更容易學的東西。

y = F(x) + x \qquad\Longrightarrow\qquad \frac{\partial y}{\partial x} = \frac{\partial F}{\partial x} + 1

對跳接做微分,會自動、無條件地產生一個 +1。

逐個符號看:x 是區塊的輸入;F(x) 是區塊的層所計算的任何東西;多出來的 +x 就是跳接/恆等連接。現在做微分,得到這個區塊的局部連鎖律因子:\partial y / \partial x = \partial F/\partial x + 1。魔法就在那個被保證的 +1。回想第 2 節:前層的梯度是這些逐層因子的乘積,而當因子很小時它會消失。但在這裡,即使 \partial F/\partial x 極小甚至幾乎為零,這個因子仍然約為 1,全靠那個 +1。把一串這樣的因子放進那個乘積,原本 0.25 \times 0.25 \times \dots 會塌縮成零,現在卻變成 (1 + \text{很小}) \times (1 + \text{很小}) \times \dots,會維持在接近 1。梯度永遠有一條乾淨、大小為 1 的路徑直通回前層——它再也無法完全消失。

一個殘差區塊:輸入 x 既穿過各層 F(x),也直接跳過去被加回來。那條跳接就是梯度搭著回家的快速車道。

殘差區塊示意圖:輸入分岔,一條路徑穿過兩個權重層計算 F(x),另一條直接跳到一個加法節點,把 F(x) 與 x 相加。

梯度裁剪:對抗梯度爆炸的安全帶

正規化與殘差大多馴服了梯度消失/爆炸問題消失那一面。但爆炸仍然會發生,尤其在循環網路、Transformer,以及任何一不小心闖進損失曲面不穩定區域的訓練中。危險來得很突然:單一個倒楣的小批次就能產生一個巨大的梯度,最佳化器邁出一大步,精心訓練的權重瞬間被炸成胡言亂語——你會看到損失在單一次迭代中飆到無限大或變成 NaN,幾小時的訓練毀於一旦。

梯度裁剪正是為這種撞車準備的安全帶。最常見的形式是範數裁剪:它把整個梯度向量的長度上限封住——如果梯度比選定的門檻還長,就把它縮回到那個門檻——但保持指向同一個方向。(一個更簡單的近親是數值裁剪,改成把每個分量各自限制在一個範圍內,較粗糙但便宜。)範數裁剪的關鍵特性是:它保留了這一步想往哪走,只限制了它走多遠

\text{if } \lVert g \rVert > c: \qquad g \;\leftarrow\; c\,\frac{g}{\lVert g \rVert}

若梯度長度超過門檻,就把它縮回那個長度,同時保持方向不變。

逐個符號看:g 是梯度向量(更新所需的所有偏導數疊在一起)。\lVert g \rVert 是它的範數——它的長度、整體大小,計算方式是各分量平方和再開根號。c選的門檻(一個超參數,常設成 1 或 5 之類)。只有在 \lVert g \rVert > c 時這個條件才會觸發。修正就是 g / \lVert g \rVert 這一項:把一個向量除以它自己的長度,會得到一個單位向量——方向相同、長度剛好是 1——再乘上 c,就把它拉伸到長度剛好是 c。所以方向完全不動,只有危險的大小被封頂。實際算一次:假設 \lVert g \rVert = 50 而你設 c = 5。由於 50 > 5,你用 5/50 = 0.1 縮放:梯度的每個分量都乘上 0.1,把一個長度 50 的向量變成長度 5、指向完全相同方向的向量。一個本來會釀災的步伐,就此變得理智。

def clip_grad_norm(grads, c):
    # grads: list of gradient arrays (one per parameter tensor)
    # Compute the global norm ||g|| across ALL of them together
    total_norm = sum((g ** 2).sum() for g in grads) ** 0.5

    if total_norm > c:
        scale = c / (total_norm + 1e-6)      # e.g. 5 / 50 = 0.1
        grads = [g * scale for g in grads]   # shorter length, SAME direction
    # if total_norm <= c, leave the gradients untouched
    return grads
範數裁剪的實作:量測全域梯度長度,只有在超過門檻 c 時才縮放。