§01學習重點
- 說出第 6 章與本章的分工:一邊決定拿到梯度之後往哪走,一邊負責把梯度生出來;並指出「反向傳播是一種最佳化演算法」錯在哪裡
- 讀懂鏈鎖法則展開式裡的每一個偏導數符號,說清楚它「釘住了什麼、放開了什麼」,不再把它跟全導數混為一談
- 在一條三站的純量產線上親手跑完一趟前向與兩趟反向,並用中央差分驗證自己沒算錯
- 寫出向量版的反向遞迴式,說出權重矩陣的轉置為什麼出現在回程、門檻遮罩的資料是誰留下來的
- 分別估算反向傳播的計算代價與記憶體代價,並說明為什麼是記憶體、而不是算力,決定了你能訓練多大的模型
- 解釋梯度為什麼隨深度呈指數而非線性變化,並區分梯度消失、梯度爆炸、學習率過大這三件常被混在一起講的事
- 從期望值與變異數的定義推到 He 初始化的「2 除以輸入維度」,並說出分子那個 2 是從哪一步冒出來的
- 講出全零初始化失效的兩個彼此獨立的理由,並各用一段程式把它演一次
§02課程內容
一、算一次梯度到底有多貴
第 6 章給了一個很乾淨的答案:想讓損失變小,就沿著讓它下降的方向走一步。這句話裡藏了一個沒交代的前提——那個方向得先有人算出來。方向的來源是梯度(gradient),也就是損失對每一個參數的偏導數(partial derivative)排成的一個向量:參數有幾個,這個向量就有幾個分量。本章要做的事,就是把這個向量算出來,而且要算得夠便宜。
兩章的分工值得先釘死,因為這是最常見的錯位之一:反向傳播(backpropagation)不是一種最佳化演算法。 它不決定步長、不決定要不要加動量、不決定不同參數該不該用不同的步幅——那些全是第 6 章的事。反向傳播只做一件事:給定目前這組參數,把損失對每一個參數的偏導數算出來,算完就交棒。把它說成「一種學習規則」,等於把「怎麼求導」跟「求完導怎麼用」壓成同一件事;等一下談到梯度爆炸的時候,你就會因為這個混淆而找錯病因。
那麼求導這件事為什麼貴?因為它要被執行的次數是三個大數字相乘。一趟訓練要走幾十萬到幾百萬個迭代步;每一步都要對全部參數各給出一個偏導數;而每一步的損失又是一批樣本各自損失的平均。把迭代步數、參數個數、批次筆數乘起來,就是整趟訓練實際要做的求導次數。當一個模型的參數多到你連印出來都不可能時,任何「一個參數配一次微分」的做法都會當場破產。
還有一個跟次數無關、更根本的理由。深層網路是一個層層套起來的複合函數。如果你把整個網路老實展開成一條式子,再對最前面那層的某個權重直接微分,得到的表達式會長得離譜——而且裡面會反覆出現同一批子項。同一個中間結果被重算了好幾遍,這就是可以省下來的東西。
所以本章要達到的效率標準是這樣的:整個網路只走兩遍——順著算一遍、倒著算一遍——就把全部參數的偏導數一次拿齊。 不是「每個參數各走一遍」,而是「總共兩遍」。做得到這件事,深層網路才變成一個訓練得動的東西。
二、把一層拆成兩個動作
先把記法備齊,後面整章都靠它。
一層網路其實做了兩個性質完全不同的動作。第一個是線性的:把上一層送來的一堆數加權相加,再加上一個常數。第二個是非線性的:把加完的結果逐個元素送進一個固定的函數。分開來看,是因為它們在求導時的行為完全不一樣——一個牽動整層所有單元,一個只管自己。
線性動作的產物叫預活化(pre-activation),非線性動作的產物叫活化值(activation)。整個網路就是這兩者交替出現:
逐項拆解:\(\mathbf{x}\) 是輸入向量,我們把它也記成 \(\mathbf{h}_0\),好讓遞迴式從 \(k=1\) 起就長得一樣。\(\boldsymbol\Omega_k\) 是第 \(k\) 層的權重矩陣,它的每一列對應這一層的一個單元、每一行對應上一層的一個單元。\(\mathbf{b}_k\) 是第 \(k\) 層的偏置向量,這一層有幾個單元,它就有幾個分量。\(\mathbf{z}_k\) 是第 \(k\) 層的預活化,\(a[\cdot]\) 是活化函數,\(\mathbf{h}_k\) 是第 \(k\) 層的活化值。全部要訓練的東西合起來記成 \(\boldsymbol\phi = \{\boldsymbol\Omega_k, \mathbf{b}_k\}\)。
假設網路一共有 \(K\) 個線性動作,那麼 \(k\) 從 1 跑到 \(K\);活化函數只加在 \(k=1\) 到 \(K-1\),最後一個線性動作的輸出 \(\mathbf{z}_K\) 就是模型的輸出,直接送進損失函數。最末層不接活化函數是慣例而不是巧合:活化函數會把輸出限制在某個值域裡(例如 ReLU 讓輸出不能為負),而模型該輸出什麼範圍的數,應該由損失函數那邊的設定決定,不該被最後一層的活化函數先斬掉一半。
兩個關於符號的提醒。
第一,索引很容易錯位。\(\mathbf{h}_k\) 是由 \(\mathbf{z}_k\) 算出來的(同號),但它被送進去的是 \(\mathbf{z}_{k+1}\)(差一號)。所以在 \(\mathbf{z}_k = \boldsymbol\Omega_k\mathbf{h}_{k-1} + \mathbf{b}_k\) 這條式子裡,等號兩邊的下標故意不一樣。後面推導時只要看到下標對不上,先回頭檢查是不是踩到這個坑。
第二,偏置在本課一律寫成 \(\mathbf{b}_k\)(粗體的羅馬字 b)。你在別的教材裡會看到有人用希臘字母寫它,本課刻意不跟——因為第 6 章已經把 \(\beta\) 指定給動量的衰減率了,同一個字母在兩章各有一義,你跨章讀的時候一定會卡住。這裡的 \(\mathbf{b}_k\) 是待訓練的參數,跟那個由你人工設定的超參數是兩種完全不同的東西。同理,權重矩陣有些材料寫成 \(\mathbf{W}_k\),本課用 \(\boldsymbol\Omega_k\),指的是同一個東西。
活化函數本章一律取修正線性單元(rectified linear unit),也就是 \(a[z]=\max(0,z)\):正的照原樣送出去,負的一律送出 0。它有兩個性質後面會反覆用到。一是它逐元素作用——第三個單元的活化值只跟第三個單元的預活化有關,不跟隔壁單元混。二是它的導數極簡單:\(z>0\) 時是 1,\(z<0\) 時是 0。至於 \(z=0\) 這一點,函數在那裡其實不可微(左右兩邊的斜率不一樣),實務上直接規定它取 0 就好——輸入剛好精確落在零點的機率極低,而且真的碰上時,取 0 或取 1 都不會讓訓練崩掉。這是工程慣例,不是數學結論,講清楚比含糊帶過好。
三、鏈鎖法則:一條三站產線的逆向追責
先把符號讀法交代清楚。\(\dfrac{\partial \ell}{\partial \omega}\) 讀作「把其他東西釘住不動,只讓 \(\omega\) 動一點點,損失 \(\ell\) 會跟著動多少」——它是一個比值:損失的變動量除以 \(\omega\) 的變動量,在變動量趨近於零時的極限。方向感很重要:分子是結果,分母是原因。
「釘住其他東西」這句話是大部分人第一次卡住的地方。因為 \(z_2\) 明明依賴 \(z_1\),如果連 \(z_2\) 也釘住,那 \(\omega_1\) 動了半天豈不是什麼都不會發生?
解法是分清楚兩層意思。鏈鎖法則(chain rule)裡的每一個因子,都是某一個單獨函數對它自己的直接輸入的偏導數。 以 \(\partial z_2/\partial h_1\) 為例:\(z_2 = \omega_2 h_1 + b_2\) 這條式子只吃 \(h_1\)、\(\omega_2\)、\(b_2\) 三個東西,這裡「釘住其他」釘的是 \(\omega_2\) 與 \(b_2\),不是釘 \(z_1\)。至於 \(h_1\) 動了之後 \(z_2\) 再怎麼往下影響 \(\ell\),那是鏈條上其他因子的工作,跟這個因子無關。
把全部因子乘起來之後,得到的 \(\partial\ell/\partial\omega_1\) 就不再是「釘住一切」的意思了——它是「讓 \(\omega_1\) 動一點,後面該跟著動的全部跟著動,最後 \(\ell\) 動多少」。這其實是一個全導數;只是因為我們這條網路是一條直線、每個中間量只有一條路往下走,寫成偏導記號不會出錯。一旦某個中間量同時被送往兩個地方(第 11 章的跳接架構與第 12 章的注意力機制都會出現這種分支),總導數就必須把每條路徑各算一次再相加,那時候偏導與全導的差別就不能再含糊。
有了這個理解,一條三站的純量網路,損失對第一站權重的導數就可以拆成六個因子:
逐項拆解:最左邊那項是損失對輸出的導數,形式由你選的損失函數決定;接下來每兩項一組,一組是「線性動作對它的輸入」的導數(結果就是那一層的權重),一組是「活化函數對它的輸入」的導數(ReLU 的話就是 0 或 1);最右邊那項是線性動作對權重自己的導數,結果是這一站收到的輸入值。因子的個數跟參數所在的深度直接對應:參數越靠前,鏈條越長,乘的項越多。
現在把它跑一次。設想一條三站的產線,每一站有一個增益係數 \(\omega_k\) 與一個校正常數 \(b_k\),站上還有一個閘門:讀數為負就當作停線、輸出 0(這就是 ReLU)。末站不設閘門,讀數直接當成品交出去。假設投料 \(x=1.5\),規格是 2.0,損失取偏差的平方;三站的參數假設為 \(\omega=(0.8,\,1.5,\,0.6)\)、\(b=(0.4,\,-0.5,\,0.3)\)。
前向一路算下去:\(z_1 = 0.8\times1.5+0.4 = 1.6\),過閘門後 \(h_1 = 1.6\);\(z_2 = 1.5\times1.6-0.5 = 1.9\),\(h_2 = 1.9\);\(z_3 = 0.6\times1.9+0.3 = 1.44\)。成品比規格少了 0.56,損失是 \(0.56^2 = 0.3136\)。
反向從末端起算。損失對 \(z_3\) 的導數是 \(2\times(1.44-2.0) = -1.12\)。這個數往回走:\(\partial\ell/\partial z_2 = -1.12\times 0.6 \times 1 = -0.672\)(乘末站增益,再乘第二站閘門開著時的 1);\(\partial\ell/\partial z_1 = -0.672\times 1.5 \times 1 = -1.008\)。有了這三個數,每一站的參數導數就只是一次乘法:偏置的導數直接等於該站的 \(\partial\ell/\partial z\),權重的導數等於 \(\partial\ell/\partial z\) 乘上那一站收到的輸入。例如 \(\partial\ell/\partial\omega_1 = -1.008\times 1.5 = -1.512\)。
注意這裡發生了什麼:三個 \(\partial\ell/\partial z\) 是接力算出來的,每一個都直接用上一個的結果,而不是各自把鏈條從頭乘一遍。這就是那個「重複子項」被省下來的地方,也是整套演算法便宜的唯一原因。
# 假設的三站產線:每站一個增益係數 w 與一個校正常數 b,
# 站上的閘門是 ReLU(讀數為負就停線,輸出 0)。末站不設閘門。
x, y = 1.5, 2.0
w = [0.8, 1.5, 0.6]
b = [0.4, -0.5, 0.3]
def run(w, b):
z1 = w[0] * x + b[0]; h1 = max(z1, 0.0)
z2 = w[1] * h1 + b[1]; h2 = max(z2, 0.0)
z3 = w[2] * h2 + b[2] # 末站直接輸出
return z1, h1, z2, h2, z3, (z3 - y) ** 2
z1, h1, z2, h2, z3, loss = run(w, b)
print(f"前向 z1={z1:.4f} h1={h1:.4f} z2={z2:.4f} h2={h2:.4f} z3={z3:.4f} 損失={loss:.4f}")
dz3 = 2 * (z3 - y) # 起點:損失對末站預活化
dw3, db3 = dz3 * h2, dz3
dz2 = dz3 * w[2] * (1.0 if z2 > 0 else 0.0)
dw2, db2 = dz2 * h1, dz2
dz1 = dz2 * w[1] * (1.0 if z1 > 0 else 0.0)
dw1, db1 = dz1 * x, dz1
hand = [dw1, db1, dw2, db2, dw3, db3]
print("手算:", " ".join(f"{v:.4f}" for v in hand))
eps, num = 1e-6, [] # 中央差分核對
for i in range(3):
for arr in (w, b):
arr[i] += eps; lp = run(w, b)[5]
arr[i] -= 2 * eps; lm = run(w, b)[5]
arr[i] += eps
num.append((lp - lm) / (2 * eps))
print("數值:", " ".join(f"{v:.4f}" for v in num))
print("最大差:", f"{max(abs(a - c) for a, c in zip(hand, num)):.2e}")實跑輸出:前向是 z1=1.6000 h1=1.6000 z2=1.9000 h2=1.9000 z3=1.4400 損失=0.3136;手算與數值兩排都是 -1.5120 -1.0080 -1.0752 -0.6720 -2.1280 -1.1200,最大差 1.21e-10。這個核對方式值得記住:中央差分(central difference)是把某個參數往上、往下各挪一個極小量,用損失的差除以挪動的總量,直接逼近導數的定義。它慢到不能拿來訓練——每個參數要跑兩次前向——但正因為它笨,它不會犯跟反向傳播一樣的錯,是驗證自己有沒有算歪最可靠的工具。
比喻: 一條三站的生產線出了瑕疵成品,要做逆向追責。你不會從第一站開始往下猜,而是從成品往回走:先量出成品偏離規格多少,再問末站「你把上游送來的半成品放大了幾倍」,把偏差按這個倍數折回給上一站;上一站再照它自己的倍數往更上游折一次。走完三站,每一站該分攤多少就都出來了,而且整趟只走了一遍——這就是鏈鎖法則由後往前算的樣子,每一段乘上去的因子就是那一站的放大倍數。這個比喻在一個地方明確失準:產線上的責任通常有個總量,分給甲的就不會再分給乙;導數不守恆,它可以在某一站被放大十倍,也可以在下一站被閘門歸零,各站加起來完全不必等於一開始那個偏差。「追責」這個詞只是用來記住「由後往前、逐站折算」這個順序,不要把守恆的直覺一起搬過來。
四、推廣到矩陣:一趟前向、兩趟反向
把純量換成向量,整套流程的骨架完全一樣,只是每一步從乘一個數變成乘一個矩陣。
前向趟(forward pass) 依序算出 \(\mathbf{z}_1, \mathbf{h}_1, \mathbf{z}_2, \mathbf{h}_2, \dots, \mathbf{z}_K\),一路算到損失。關鍵動作是:這些中間量算完不要丟,全部留在原地。 因為反向趟會用到它們,而且用途分成兩種——一種只需要 \(\mathbf{z}\) 的正負號(用來判斷閘門開著還是關著),另一種需要 \(\mathbf{h}\) 的實際數值(用來當乘法的另一個因子)。
回程要算的東西分成兩類,習慣上叫「兩趟」。先講清楚免得誤會:這裡的兩趟不是把網路倒著走兩遍,而是在同一次由後往前的走訪裡,每退到一層就交錯做兩件事——先算損失對這一層中間量的導數,再用它算出這一層兩組參數的導數。整趟下來網路仍然只被倒著走了一遍。
反向趟的第一趟算的是損失對各層中間量的導數。起點是損失對最末預活化的導數 \(\mathbf{g}_K = \partial\ell/\partial\mathbf{z}_K\),它長什麼樣由損失函數決定;接著往前遞推:
逐項拆解:\(\mathbf{g}_k\) 是損失對第 \(k\) 層預活化的導數向量。\(\boldsymbol\Omega_k^{\mathsf{T}}\) 是權重矩陣的轉置(transpose),把它左乘上去,就完成了「從第 \(k\) 層退回第 \(k-1\) 層」這一步。\(\mathbb{I}[\mathbf{z}_{k-1}>0]\) 是一個由 0 與 1 組成的遮罩向量:對應的預活化為正就是 1,否則是 0。\(\odot\) 表示逐元素乘法(elementwise multiplication),也就是兩個等長向量對應位置各自相乘。
兩件事值得多說一句。
轉置為什麼會出現在回程? 一個線性映射對它輸入的導數,就是這個映射矩陣的轉置。完整的矩陣微積分推導超出本課範圍,但你可以用形狀自己驗一遍,而且應該養成這個習慣。設第 \(k\) 層有 \(D_k\) 個單元,那麼 \(\boldsymbol\Omega_k\) 是 \(D_k \times D_{k-1}\)、\(\mathbf{g}_k\) 是 \(D_k \times 1\)。要得到一個 \(D_{k-1}\times 1\) 的結果,唯一能對接的乘法就是 \(D_{k-1}\times D_k\) 乘 \(D_k\times 1\)——而 \(D_{k-1}\times D_k\) 正是 \(\boldsymbol\Omega_k\) 轉置後的形狀。矩陣導數的排列慣例是很多人第一次寫錯的地方,形狀相容性是最省事的自我校驗。
遮罩為什麼是逐元素乘,而不是矩陣乘? 嚴格寫的話,活化函數對它輸入的導數是一個矩陣。但因為 ReLU 逐元素作用、第 \(i\) 個單元的活化值只依賴第 \(i\) 個預活化,這個矩陣的非對角位置全是 0——它是一個對角矩陣。而「乘上一個對角矩陣」跟「把對角線上那串數字拿去逐元素相乘」結果完全一樣,後者省掉了一整個矩陣的儲存與大部分乘法。這個替換不是近似,是恆等。
反向趟的第二趟才輪到參數。有了 \(\mathbf{g}_k\),這一層的兩組導數各只要一步:
逐項拆解:偏置的導數就等於該層預活化的導數,因為在 \(\mathbf{z}_k = \boldsymbol\Omega_k\mathbf{h}_{k-1}+\mathbf{b}_k\) 裡,偏置的係數是 1,這一步的導數是 1,乘上去等於沒乘。權重的導數是一個外積(outer product):一個 \(D_k\times1\) 的向量乘上一個 \(1\times D_{k-1}\) 的向量,得到 \(D_k\times D_{k-1}\) 的矩陣——形狀跟 \(\boldsymbol\Omega_k\) 一模一樣,這是它對的必要條件。外積的意思很直白:權重矩陣第 \(i\) 列第 \(j\) 行那個數的導數,等於「第 \(i\) 個單元收到的導數」乘上「第 \(j\) 個輸入的值」。一個權重對損失的影響力,正比於它當時實際乘到的那個數有多大——這也解釋了為什麼前向趟非把 \(\mathbf{h}\) 存下來不可。最前面那一層是同一條式子的特例:那裡的 \(\mathbf{h}_0\) 就是輸入 \(\mathbf{x}\) 本身。
以上都是單一樣本的導數。實際訓練時損失是整批樣本的平均,\(\mathcal{L}=\frac{1}{I}\sum_{i=1}^{I}\ell_i\),所以每個參數的梯度就是各樣本導數的平均。(有些教材把損失寫成總和、不除以 \(I\),兩者只差一個正的常數倍,最小值的位置一樣。)這個平均後的向量,就是交給第 6 章那些更新規則的東西。
底下這段程式把上面的遞迴式原樣寫出來,並用中央差分逐個參數核對:
import numpy as np
rng = np.random.default_rng(0)
dims = [4, 6, 6, 6, 1] # 輸入 4 維 → 三層寬 6 → 輸出 1 維
K = len(dims) - 1 # 線性動作的個數
Om = [rng.normal(0, 0.7, (dims[k + 1], dims[k])) for k in range(K)]
bs = [rng.normal(0, 0.1, (dims[k + 1], 1)) for k in range(K)]
x, y = rng.normal(0, 1, (dims[0], 1)), 0.5
def forward(Om, bs): # 前向趟:把 z 與 h 全部存下來
h, z = [x], []
for k in range(K):
zk = Om[k] @ h[-1] + bs[k]
z.append(zk)
h.append(np.maximum(zk, 0.0) if k < K - 1 else zk)
return z, h, float((h[-1][0, 0] - y) ** 2)
def backward(Om, bs):
z, h, loss = forward(Om, bs)
g = 2.0 * (h[-1] - y) # 起點:損失對最末預活化
dOm, dbs = [None] * K, [None] * K
for k in range(K - 1, -1, -1):
dbs[k] = g # 偏置導數=該層預活化導數
dOm[k] = g @ h[k].T # 權重導數=外積
if k > 0:
g = (Om[k].T @ g) * (z[k - 1] > 0) # 轉置左乘 + 門檻遮罩
return dOm, dbs, loss
dOm, dbs, loss = backward(Om, bs)
eps, worst = 1e-6, 0.0
for k in range(K):
for arr, ana in ((Om[k], dOm[k]), (bs[k], dbs[k])):
for idx in np.ndindex(arr.shape):
arr[idx] += eps; lp = forward(Om, bs)[2]
arr[idx] -= 2 * eps; lm = forward(Om, bs)[2]
arr[idx] += eps
worst = max(worst, abs((lp - lm) / (2 * eps) - ana[idx]))
print(f"損失={loss:.6f} 參數個數={sum(a.size for a in Om + bs)}")
print("dOmega 各層形狀:", [a.shape for a in dOm])
print(f"反向傳播 vs 數值微分,最大絕對差 = {worst:.3e}")實跑輸出:損失=0.005154,參數個數=121,各層權重導數的形狀是 [(6, 4), (6, 6), (6, 6), (1, 6)](跟權重本身一致),反向傳播與數值微分的最大絕對差是 1.402e-11。注意兩者的成本差距:反向傳播跑一趟前向加一趟反向就拿到 121 個導數;中央差分為了核對這 121 個數,跑了 242 次前向。參數越多,這個差距越誇張。
兩種代價,分開算。 先看算力。把兩趟裡的浮點運算逐項數過去,你會發現壓倒性的多數都落在同一件事上:矩陣與向量相乘。單看一次乘法或一次加法,硬體做起來便宜到可以忽略;真正堆出成本的是次數,而次數由相鄰兩層的寬度(再乘上這一批有幾筆樣本)決定。所以估算算力時,你要盯的是矩陣的形狀,不是式子裡出現了幾個看起來很嚇人的符號。記憶體代價則是另一回事:批次裡每一筆樣本的每一層中間量都得同時待在記憶體裡,直到反向趟走到那一層為止。所以記憶體用量正比於「批次大小 × 深度 × 每層寬度」,而參數本身的記憶體跟批次大小無關。訓練大模型時真正先撞牆的通常是記憶體而不是算力,作業三會讓你把這兩個數字實際算出來比一比。這兩者之間可以交換:只留幾個關卡的紀錄、其餘等反向趟走到時再重算一次前向,就能把記憶體換成算力,這個做法叫梯度檢查點(gradient checkpointing),本課只點到它存在。
框架替你做了什麼。 你不必每換一個架構就重推一次上面的式子——這件事有名字,叫自動微分(automatic differentiation)。
先講方向的選擇,因為那才是重點。鏈鎖法則要把一長串局部導數乘起來,而乘法可以從任一端開始結合:從輸入那端往輸出推,叫前向模式(forward mode);從輸出那端往輸入推,叫反向模式(reverse mode)。兩種算出來的答案完全一樣,成本卻差得離譜。訓練時你要的是一個純量損失對海量參數的導數:從那一個數往回發散,走一趟就把所有參數的導數一次收齊;反過來從參數那端往前推,有幾個參數就得推幾趟。這個不對稱就是深度學習只用反向模式的全部理由,也正是你前面手寫那兩趟在做的事。
至於「自動」自動在哪裡:每種運算元件只需要知道一件非常局部的事——給我輸入,我交代得出輸出對輸入的導數長什麼樣。這一小塊知識寫一次就永遠有效,跟它被放在哪個網路裡無關。剩下的交給框架:你搭網路時,框架順手把運算的相依關係記成一張計算圖(computation graph),之後照這張圖的順序跑前向、照反過來的順序把各元件的局部導數串起來。所以你怎麼接都不必重推,也不限於本章這種一條龍的接法:中間分岔出去、後面再匯回來都算數,唯一撐不住的是讓某個值繞一圈回到它自己身上。本站是零依賴的環境,只用 numpy,所以不做框架示範;知道它在做的就是你剛剛手寫的那兩趟,就夠了。
比喻: 回到那條產線,這次談的是逆向追責的前置條件。要往回追,你得調得出每一站當時的在製品狀態:那一站的讀數是正是負(閘門開著還是關著)、送出去的半成品是多少。這些不是追責當下才去量的,是生產當時就順手記下來的——沒記,事後就只能把整批料重跑一遍。這就是前向趟為什麼要把中間量留在原地。失準之處在於:真實產線把紀錄弄丟了,重做要付出實打實的材料成本;而在網路裡,中間量弄丟了確實可以拿同一組參數把前向趟重算一次,代價只有算力,沒有材料。所以這裡不是「有沒有紀錄」的非黑即白,而是一條可以連續調節的取捨——上一段提到的梯度檢查點,就是刻意只記幾站、其餘用時再算。
五、爆炸與消失:逐站增益的連乘
到這裡演算法已經齊了。接下來的問題是:參數的起點該怎麼設? 這個問題聽起來像細節,實際上決定了深層網路動不動得了。
先看前向趟。每過一層,訊號的幅度會被放大或縮小一點點——放大多少,取決於那一層權重的尺度。如果每層平均把幅度乘上 0.9,走三十層之後就只剩下原本的 \(0.9^{30}\);如果每層乘上 1.1,三十層之後會是 \(1.1^{30}\)。這兩個數差多少?
這就是「指數級」三個字的全部含義:逐層的效果是連乘,不是累加。 如果每層的效果是相加,走三十層之後兩邊也不過相差 \(30\times0.2=6\),那是線性的、看得住的;連乘卻讓「每層只差 0.2」變成最後差了四百倍以上。而且層數越多,同一個微小偏差被放大得越厲害——它出現在指數的位置上。這一點是本章最需要記住的直覺,也是很多人明明看得懂公式卻仍然低估初始化重要性的原因。
反向趟有一模一樣的體質。回頭看第四小節那條遞迴式:每往回退一層,就要左乘一次 \(\boldsymbol\Omega_k^{\mathsf{T}}\)。退三十層就是三十個轉置矩陣連乘。所以前向的幅度往哪個方向跑,反向的梯度大致也往同一個方向跑。
兩個方向各自有名字。梯度縮到極小、更新量小到跟沒更新一樣,叫梯度消失(vanishing gradient);梯度膨脹到極大、一步就把參數扔到很遠的地方,叫梯度爆炸(exploding gradient)。極端情況下浮點數本身也會出事:太小會下溢成 0,太大會上溢成無限大,那時候連數字都不剩了。
底下用程式把這件事跑出來。設一個每層 64 個單元、疊 30 層的網路(偏置全設 0),權重的變異數分別取 \(1/D\)、\(2/D\)、\(4/D\),其中 \(D\) 是每層的單元數:
import numpy as np
D, L = 64, 30 # 每層 64 個單元,疊 30 層,偏置全設 0
def sweep(c): # 權重變異數取 c/D
rng = np.random.default_rng(0)
Om = [rng.normal(0, np.sqrt(c / D), (D, D)) for _ in range(L)]
h = rng.normal(0, 1, (D, 1))
masks, fwd = [], [float(np.sqrt(np.mean(h ** 2)))] # 第 0 層=輸入
for W in Om: # 前向:記下每層活化值的均方根
z = W @ h
m = z > 0
h = z * m
masks.append(m)
fwd.append(float(np.sqrt(np.mean(h ** 2))))
g, bwd = rng.normal(0, 1, (D, 1)), []
for W, m in zip(reversed(Om), reversed(masks)): # 反向:同一批權重的轉置
g = W.T @ (g * m)
bwd.append(float(np.sqrt(np.mean(g ** 2))))
return fwd, bwd
for c, name in ((1.0, "1/D"), (2.0, "2/D"), (4.0, "4/D")):
f, b = sweep(c)
print(f"變異數 {name}|活化RMS 第0層 {f[0]:.3f} 第15層 {f[15]:.3e} 第30層 {f[30]:.3e}"
f"|梯度RMS 退到第1層 {b[-1]:.3e}")
print(f"假設每站增益固定為 1.1 或 0.9,30 站之後:{1.1 ** 30:.2f} 倍 對 {0.9 ** 30:.5f} 倍")實跑輸出三行。變異數取 1/D 時,活化值的均方根從第 0 層的 1.019 掉到第 15 層的 1.315e-02、第 30 層的 4.076e-05,梯度退到第 1 層時只剩 1.637e-05。取 2/D 時,三個數是 1.019、2.381e+00、1.336e+00,梯度是 5.364e-01——大致守住在 1 附近。取 4/D 時則一路衝到 4.309e+02 與 4.376e+04,梯度是 1.758e+04。最後一行印出 17.45 倍 對 0.04239 倍。
三個誤解在這裡一次拆掉。
「梯度消失就是梯度變成零」——不是。 上面那個 4.076e-05 不是零,它是一個完全合法的浮點數。問題在於它乘上學習率之後,對參數造成的改變小到被後面幾層的更新完全蓋過,於是最前面那幾層實質上不動。層數再多一點才真的會下溢成零。把它想成「零」的壞處是:你會去檢查梯度是不是剛好等於零,然後看到一堆非零的小數就放心了。真正該看的是各層梯度幅度之間的比例——最前面幾層比最後幾層小了幾個數量級。
「梯度爆炸是學習率設太大造成的」——這是兩件不同的事。 學習率太大,是梯度算得好好的、但你在那個方向上走太遠。梯度爆炸則是梯度本身在回程被逐層放大到不合理的量級,這時候你就算把學習率調到很小,各層之間的更新幅度仍然嚴重失衡——最前面幾層幾乎不動,最後幾層一步跳到天邊。兩者的處方也不同:前者調第 6 章的步長,後者調本章的起點。
「初始化會改變損失地景」——不會。 損失函數由模型結構、資料與損失定義決定,跟你從哪一組參數出發完全無關;地形是固定的。初始化只決定你站在哪一個位置。但這不代表它不重要:站的位置決定了你腳下的坡度長什麼樣,而如果坡度小到量不出來,或大到讓你一步摔出地圖,你就走不動。地形沒變,能不能走得動變了。
六、He 初始化:把那個 2 推出來
上一節用實驗看出 \(2/D\) 這個尺度剛好守得住。這一節把它算出來——這是本章唯一一段完整的推導,值得慢慢走。
前提設定。 偏置全部設成 0;每一個權重各自獨立地從一個平均值為 0、變異數為 \(\sigma^2\) 的常態分布抽出來;\(\sigma^2\) 就是我們要解的未知數。目標量是相鄰兩層預活化的變異數之比。以下把第 \(k\) 層第 \(i\) 個單元的預活化寫成 \(z_{ki} = \sum_{j=1}^{D_{k-1}}\Omega_{kij}\,h_{k-1,j}\)。
步驟 1:期望值對加總是線性的,所以 \(\mathbb{E}[z_{ki}] = \sum_j \mathbb{E}[\Omega_{kij}h_{k-1,j}]\)。(期望值就是「取很多次的平均會趨近的那個數」。)
步驟 2:\(\Omega_{kij}\) 是這一層的權重,\(h_{k-1,j}\) 只由前面各層的權重決定,兩者互相獨立,所以乘積的期望等於期望的乘積:\(\mathbb{E}[\Omega_{kij}]\,\mathbb{E}[h_{k-1,j}]\)。
步驟 3:權重的平均值是 0,於是整個和是 0,得 \(\mathbb{E}[z_{ki}]=0\)。
步驟 4:搬出變異數恆等式 \(\mathrm{Var}[u]=\mathbb{E}[u^2]-\mathbb{E}[u]^2\)。
步驟 5:既然 \(\mathbb{E}[z_{ki}]=0\),第二項消失,於是 \(\mathrm{Var}[z_{ki}]=\mathbb{E}[z_{ki}^2]\)。求變異數變成求二階動差(也就是平方的期望值),這一步是整段推導的關鍵簡化。
步驟 6:把平方展開成雙重求和 \(\sum_j\sum_{j'}\mathbb{E}[\Omega_{kij}\Omega_{kij'}h_{k-1,j}h_{k-1,j'}]\)。當 \(j\neq j'\) 時,兩個權重互相獨立而且各自平均為 0,這一項的期望是 0。所有交叉項就這樣消失了。
步驟 7:只剩下對角項,一共 \(D_{k-1}\) 項——項數就是來源層的單元數。
步驟 8:單項再用一次獨立性拆開:\(\mathbb{E}[\Omega_{kij}^2]\,\mathbb{E}[h_{k-1,j}^2] = \sigma^2\,\mathbb{E}[h_{k-1}^2]\)。合起來得
步驟 9:現在處理 \(\mathbb{E}[h^2]\)。因為偏置是 0、權重的分布左右對稱,\(z_{k-1}\) 的分布也會對稱地落在 0 兩側;而 ReLU 把負的那一半全部截成 0。
步驟 10:所以 \(\mathbb{E}[h_{k-1}^2]=\tfrac12\,\mathbb{E}[z_{k-1}^2]=\tfrac12\,\mathrm{Var}[z_{k-1}]\)。被截掉的那一半對平方和沒有貢獻,剩下那一半原封不動。
步驟 11:把步驟 8 與步驟 10 合起來:
步驟 12:我們要的是逐層守恆——變異數走過一層之後不放大也不縮小,也就是要求上式的係數等於 1。
步驟 13:解出來就是
分子那個 2,來源就是步驟 9 到 10 那個二分之一的倒數。 ReLU 砍掉一半,你就得在權重的變異數上補回兩倍。知道這件事,你才會明白換一個活化函數時這個數字為什麼要跟著換(作業二會讓你算一次)。
推導裡有兩處是假設而非定理,說清楚比較誠實:步驟 8 把同一層各單元的 \(\mathbb{E}[h^2]\) 當成同一個值,步驟 9 假設預活化的分布左右對稱。兩者在實務上都夠準,但也解釋了上一節那條 2/D 的曲線為什麼會小幅上下擺動——那是有限寬度下的隨機起伏,不是推導錯了。
這個結果叫 He 初始化(He initialization):權重的變異數取 2 除以輸入維度,這裡的輸入維度指的是這個權重矩陣所作用的那一層的寬度,也就是 \(D_{k-1}\),習慣上叫扇入(fan-in)。
它有兩個變體要知道。
第一,反向趟給出的是另一個條件。 上面守的是前向的變異數。如果改成要求梯度在回程也逐層守恆,同一套推導會給出 \(\sigma^2 = 2/D_k\)。
比較這兩條式子時,最快的方式是看分母指的是矩陣的哪一邊。\(\boldsymbol\Omega_k\) 是個 \(D_k \times D_{k-1}\) 的矩陣:它的行數 \(D_{k-1}\) 是每個單元一次要吃進多少個數,也就是上面剛講的扇入;它的列數 \(D_k\) 是這個矩陣一次要產出多少個數,叫扇出(fan-out)。前向條件盯行數,反向條件盯列數。
於是有沒有辦法兩邊都守,就變成一個純粹的形狀問題:\(2/D_{k-1}\) 與 \(2/D_k\) 要相等,只有在矩陣是方陣、也就是相鄰兩層一樣寬的時候才成立。層寬一有變化,你就得挑一邊守、另一邊放掉。實務上多半兩邊都不放掉,改成折衷:把 \(2/D\) 裡的 \(D\) 換成兩層寬度的平均 \((D_{k-1}+D_k)/2\),得到 \(\sigma^2 = 4/(D_{k-1}+D_k)\)。分子從 2 變成 4 沒有別的玄機,就是分母除以 2 的結果。
第二,你會遇到另一個少了一倍的版本。 Xavier 初始化(Xavier initialization,也叫 Glorot 初始化)的變異數比 He 小一半。差別的來源就是步驟 9 到 10:Xavier 的推導針對的是像 sigmoid、tanh 這類在原點附近大致對稱、不會把一半的值砍掉的活化函數,所以沒有那個二分之一要補。實務上的判準很單純:用 ReLU 系的活化函數就用 He,用對稱型的就用 Xavier。
偏置一律設成 0。 這是慣例,理由是偏置不承擔打破對稱的責任——真正需要彼此不同的是權重。至於為什麼一定要有東西「彼此不同」,就是下面這件事。
全零初始化為什麼不行? 這裡有兩個彼此獨立的理由,兩個都要知道,因為堵住其中一個並不能救另一個。
理由一:對稱性永遠破不了。 如果同一層各單元收到的權重都相同,它們在前向趟會算出一模一樣的值;如果下一層送回給它們的權重也都相同,它們在反向趟就會收到一模一樣的導數,於是更新之後仍然一模一樣。(兩個條件缺一不可,而全零或全同值的起點剛好兩個都滿足。)假設這一層有一百個單元,它們永遠是一百份同樣的東西,實質上等於只有一個單元。這條理由跟「零」無關,跟「全都一樣」有關——把權重全設成 0.3 也照樣中招。
理由二:更新量本身就是零。 權重的導數是 \(\mathbf{g}_k\mathbf{h}_{k-1}^{\mathsf{T}}\)。全零起點下,前向趟算出來的 \(\mathbf{h}\) 全是 0,這個外積整個是零矩陣;同時回程要左乘的 \(\boldsymbol\Omega^{\mathsf{T}}\) 也全是 0,梯度在第一步就被截斷。這條理由專屬於「零」,換成任何非零常數就不成立。
import numpy as np
x, y = np.array([[1.0], [-0.5], [2.0]]), 1.0 # 3 維輸入、4 個隱藏單元、1 維輸出
def grads(O1, b1, O2, b2):
z1 = O1 @ x + b1
h1 = np.maximum(z1, 0.0)
z2 = O2 @ h1 + b2
g2 = 2.0 * (z2 - y)
g1 = (O2.T @ g2) * (z1 > 0)
return g1 @ x.T, g1, g2 @ h1.T, g2 # dO1, db1, dO2, db2
zeros = (np.zeros((4, 3)), np.zeros((4, 1)), np.zeros((1, 4)), np.zeros((1, 1)))
dO1, db1, dO2, db2 = grads(*zeros)
print("全零起點:dOmega1 全為零?", np.allclose(dO1, 0),
" dOmega2 全為零?", np.allclose(dO2, 0),
" db2 =", float(db2[0, 0]))
const = (np.full((4, 3), 0.3), np.full((4, 1), 0.3),
np.full((1, 4), 0.3), np.full((1, 1), 0.3))
print("同值起點:dOmega1 的四列")
for r in grads(*const)[0]:
print(" ", np.round(r, 6))
rng = np.random.default_rng(0)
he = (rng.normal(0, np.sqrt(2 / 3), (4, 3)), np.zeros((4, 1)),
rng.normal(0, np.sqrt(2 / 4), (1, 4)), np.zeros((1, 1)))
print("He 起點:dOmega1 的四列")
for r in grads(*he)[0]:
print(" ", np.round(r, 6))實跑輸出:全零起點下 dOmega1 全為零? True、dOmega2 全為零? True,整個網路只有末層偏置拿到 db2 = -2.0,其餘參數一動也不動。同值起點(全部設成 0.3)下,導數不再是零了,但四列一字不差,全是 [0.336 -0.168 0.672]——理由二解掉了,理由一還在。He 起點下四列分別是 [10.243154 -5.121577 20.486309]、[0.963908 -0.481954 1.927817] 以及兩列全零;那兩列之所以是零,是因為那兩個單元對這一筆輸入剛好落在閘門的負區,這是 ReLU 的正常現象,而且正好證明了各單元的命運已經不同——對稱性確實破掉了。
把這一章收成一句話:梯度不是憑空出現的,它是一趟前向加兩趟反向算出來的;而它算出來之後有多大,在你設定第一組參數的那一刻就已經大致決定了。
§03原書對照
原書第 7 章把兩件事併在一起處理:梯度怎麼算得快,以及參數該從哪裡起跳。本課重排了敘事順序,也換掉了原書用來示範的那個複合函數例子,因此有幾處原書鋪陳得更完整,值得進階讀者翻回去對照。
第一處是原書 7.3 節的玩具示範(pp.100–103)。原書刻意挑了一個由正弦、指數、餘弦三種函數層層套起來的八參數純量模型:先在 p.100 攤開直接對最前面那個權重微分會得到多長的一條式子,讓讀者親眼看見裡頭重複出現的子項;再用 pp.101–103 三頁把前向、反向第一趟、反向第二趟逐條寫完,配三張流程圖標出每個中間量存在哪裡、又在哪一步被取用。本課另造了一個等效的手算示範,想看純量版的完整逐行推導,那三頁最省力。
第二處是矩陣微積分的細節。本課直接給出「回程要乘上權重矩陣的轉置」這個結論,原書則在 p.105 把它列成獨立一式,並把證明交給附錄 B.5 與該章的兩道進階習題;同一頁也用一樣的方式處理權重導數的外積形式。p.104 另標出鏈鎖式裡每一項的實際尺寸,對想動手實作的人特別有用。緊接著 p.106 把整套演算法壓成一組遞迴式,並點出它算得快但吃記憶體的體質。
第三處是初始化的完整推導(pp.108–111)。原書從期望值與變異數的定義一路推到結論,逐行標明每一步用到哪一條期望值規則、在哪裡假設了權重與活化值互相獨立,整段寫得像一張可以直接拿來演算的稿紙。p.110 的圖用五十層、每層一百個單元的網路實測五組不同的權重變異數,把前向活化與反向梯度的幅度隨層數的變化畫在同一張對數座標上,是全章最直觀的一頁。權重矩陣不是方陣時的折衷式子在 p.111。
第四處是把手算搬進框架的那一段。原書 p.107 說明現代框架怎麼讓每個運算元件各自負責自己那一小塊導數,再由框架掌握整串運算的順序,自動組出前後兩趟;同一頁也交代了整批樣本一起平行處理之後,輸入為什麼會從向量升格成張量,以及影像這類資料會讓維度數再多出幾層。原書在 p.107 也把這套做法的適用範圍講明白了:分支與匯流都不影響,只要計算圖沒有環,前後兩趟就都成立。
第五處是延伸方向。原書 pp.113–114 的註記交代了反向傳播被反覆發現的歷史、兩種常見初始化為何差一個倍數,以及好幾種本課沒提的初始化方案;同一段還介紹了用重算換記憶體的兩種手法,以及三類分散式訓練策略。習題方面,pp.115–117 有一題要你嚴格證明那個二分之一係數,有一題問全零初始化會發生什麼事,另外兩題讓你用兩種相反方向各算一次同一個計算圖,親自體會為什麼深度網路只用其中一種。p.112 另附一份完整的訓練範例碼。原書第 7 章對應印刷頁 pp.96–117。
§04作業和解答
作業一:一條會停線的四站產線
把第三小節的示範延長成四站。假設投料 \(x=2.0\)、規格 \(y=1.0\),四站的參數為 \(\omega=(0.5,\,-1.4,\,0.9,\,0.7)\)、\(b=(0.2,\,0.3,\,0.1,\,-0.2)\),前三站有 ReLU 閘門、末站沒有,損失仍取偏差的平方。(a)算出全部前向量與損失。(b)算出八個參數的導數。(c)\(\omega_1\)、\(b_1\)、\(\omega_2\)、\(b_2\) 這四個為什麼全是零?這對訓練代表什麼?(d)\(\omega_3\) 是零而 \(b_3\) 不是零,為什麼?
解答 SOLUTION
(a)\(z_1 = 0.5\times2.0+0.2 = 1.2\),\(h_1 = 1.2\);\(z_2 = -1.4\times1.2+0.3 = -1.38\),閘門關閉,\(h_2 = 0\);\(z_3 = 0.9\times0+0.1 = 0.1\),\(h_3 = 0.1\);\(z_4 = 0.7\times0.1-0.2 = -0.13\)。損失 \(=(-0.13-1.0)^2 = 1.2769\)。
(b)由後往前:\(\partial\ell/\partial z_4 = 2\times(-1.13) = -2.26\),於是 \(\partial\ell/\partial b_4 = -2.26\)、\(\partial\ell/\partial\omega_4 = -2.26\times h_3 = -0.226\)。退一層:\(\partial\ell/\partial z_3 = -2.26\times0.7\times1 = -1.582\),所以 \(\partial\ell/\partial b_3 = -1.582\)、\(\partial\ell/\partial\omega_3 = -1.582\times h_2 = 0\)。再退一層時遇到關閉的閘門:\(\partial\ell/\partial z_2 = -1.582\times0.9\times0 = 0\),因此 \(\partial\ell/\partial\omega_2 = \partial\ell/\partial b_2 = 0\),而且 \(\partial\ell/\partial z_1\) 也是 0,\(\partial\ell/\partial\omega_1 = \partial\ell/\partial b_1 = 0\)。八個導數依序為 0、0、0、0、0、-1.582、-0.226、-2.26(順序為 \(\omega_1, b_1, \omega_2, b_2, \omega_3, b_3, \omega_4, b_4\)),已用中央差分核對,兩排完全相符。
(c)因為第二站的閘門關著。ReLU 在負區的導數是 0,回程乘到 0 就整條鏈斷了,斷點以前的所有參數這一筆樣本都拿不到任何訊息。這對訓練代表:對這一筆樣本而言,前兩站等於不存在。但這不必然是災難——換一筆樣本,\(z_2\) 可能就是正的,閘門重新打開。真正的麻煩是某個單元對所有樣本都關著,那它就永遠收不到梯度、永遠不會更新,這種單元有個俗稱叫「死掉的 ReLU」。用 He 初始化、避免權重整體偏負,就是為了降低這種事發生的機率。
(d)兩者的公式不同。權重的導數是 \(\partial\ell/\partial z_3\) 乘上這一站收到的輸入 \(h_2\),而 \(h_2=0\),所以整項是零——不是因為梯度斷了,而是因為乘上了零。偏置的導數則直接等於 \(\partial\ell/\partial z_3 = -1.582\),它不乘任何輸入。這正是第六小節「全零初始化」理由二的縮影:全零起點下所有的 \(h\) 都是零,於是所有權重的導數都是零,但末層偏置照樣拿得到值。
作業二:換一個活化函數,那個 2 要變成多少
把 ReLU 換成洩漏型修正線性單元(leaky rectified linear unit):\(a[z]=z\)(當 \(z>0\))、\(a[z]=\alpha z\)(當 \(z\le 0\)),其中 \(\alpha\) 是一個介於 0 與 1 之間的小常數。(a)重做第六小節的步驟 9 到 13,求出讓變異數逐層守恆的 \(\sigma^2\)。(b)用 \(\alpha=0.1\) 與 \(\alpha=0.5\) 各檢查一次你的答案,說明為什麼其中一組看起來「差別不大」。
解答 SOLUTION
(a)只有步驟 9 到 10 需要改。負的那一半不再被截成零,而是被乘上 \(\alpha\),平方之後就是乘上 \(\alpha^2\)。所以
代回步驟 11,逐層的變異數增益變成 \(\frac{1+\alpha^2}{2}D_{k-1}\sigma^2\),令它等於 1 得
兩個端點可以拿來檢查這個式子合不合理:\(\alpha=0\) 時它退回 ReLU 的 \(2/D\);\(\alpha=1\) 時活化函數變成恆等映射(等於沒有非線性),它給出 \(1/D\),正好是不打折的守恆條件。
(b)把第五小節 sweep 函式裡的 h = z * m 換成 h = np.where(z > 0, z, alpha * z),其餘一行不動(只看前向那一段的輸出;反向那一段的遮罩要一併改成 \(\alpha\) 才算完整,本題不涉及),\(D=64\)、\(L=30\)、seed=0 下實跑。\(\alpha=0.1\) 時理論值是 \(1.9802/D\):用 2/D 跑,第 30 層活化均方根是 1.338;用 1.9802/D 跑,是 1.153——兩者都在 1 附近,肉眼分不出差別。\(\alpha=0.5\) 時理論值是 \(1.6/D\):用 2/D 跑,第 15 層已經是 5.513、第 30 層漲到 19.863;改用 1.6/D,第 15 層是 1.034、第 30 層是 0.699,穩住了。
差別為什麼一組明顯、一組不明顯?因為修正項是 \((1+\alpha^2)\)。\(\alpha=0.1\) 時它只有 1.01:每層的變異數差 1%,三十層連乘是 \(1.01^{30}\approx1.35\) 倍,換算成均方根只差 1.16 倍——實測的 \(1.338/1.153=1.16\) 正好對上,而這個幅度淹在隨機起伏裡。\(\alpha=0.5\) 時它是 1.25:三十層連乘是 \(1.25^{30}\approx808\) 倍,均方根差 28.4 倍——實測的 \(19.863/0.699=28.4\) 也對上了。這題要帶走的結論是:係數本身要算對,但它偏離多少才會出事,取決於它被連乘幾次。
作業三:算算看到底是誰先撞牆
假設一個 40 層的全連接網路,為簡化把每一層(含輸入與輸出)都當成 1024 維,中間量與參數都用 32 位元浮點數儲存(每個數 4 個位元組)。前向趟每層要存下 \(\mathbf{z}_k\) 與 \(\mathbf{h}_k\) 兩個向量。(a)批次大小 256 時,全部中間量要多少記憶體?(b)參數本身要多少記憶體?(c)批次要大到多少,中間量才會超過參數?(d)如果只存 \(\mathbf{z}_k\)、需要 \(\mathbf{h}_k\) 時當場用 ReLU 重算,(a)的答案變多少?這說明了什麼?
解答 SOLUTION
(a)每筆樣本每層存兩個 1024 維向量,就是 \(40\times2\times1024 = 81{,}920\) 個數。乘上批次 256 得 20,971,520 個數,再乘 4 個位元組是 83,886,080 個位元組,剛好 80.0 MiB。
(b)每層權重矩陣 \(1024\times1024\)、偏置 1024 個,四十層合計 \(40\times(1024\times1024+1024) = 41{,}984{,}000\) 個數,乘 4 得 167,936,000 個位元組,約 160.2 MiB。所以在批次 256 這個設定下,參數還比中間量佔得多。
(c)把兩邊設成相等:\(81{,}920\times B = 41{,}984{,}000\),解得 \(B = 512.5\),所以批次到 513 就換中間量佔上風。這裡的關鍵不在這個特定數字,而在兩者的成長方式不同——參數的記憶體跟批次無關,中間量的記憶體正比於批次。批次一大,撞牆的一定是中間量。(以上四個數字均用程式重算核對過。)
(d)\(\mathbf{h}_k = \max(0, \mathbf{z}_k)\) 可以由 \(\mathbf{z}_k\) 在需要時當場算出來,所以中間量減半,變成 40.0 MiB。這正是「用算力換記憶體」最便宜的一個例子:多花的計算量幾乎可以忽略(一次逐元素取最大值),省下的記憶體卻是一半。把同樣的思路推到極致——只留幾層的紀錄、其餘整段重跑前向——就是第四小節提到的梯度檢查點。記憶體與算力之間有一整條可以調節的取捨曲線,不是二選一。
§05參考資料
- CS231n:反向傳播與計算圖筆記 — 史丹佛課程講義,用大量小型計算圖把「局部導數怎麼串起來」演到極細,補足本章第三小節
- Calculus on Computational Graphs: Backpropagation — 用圖論的角度解釋前向模式與反向模式為什麼成本差這麼多,本章第四小節那段話的完整版
- Delving Deep into Rectifiers(He 等人,2015) — He 初始化的原始論文,第六小節那條變異數推導的出處
- Understanding the difficulty of training deep feedforward neural networks(Glorot 與 Bengio,2010) — Xavier 初始化的原始論文,拿來跟上一筆對照就看得出那個 2 倍差是怎麼來的
- Automatic Differentiation in Machine Learning: a Survey — 自動微分的完整分類與歷史,想知道框架底下究竟有幾種做法時看這篇
- NumPy:線性代數常式 — 本章所有矩陣運算的語法出處,轉置、矩陣乘、外積都在這裡
- Understanding Deep Learning(MIT Press) — 本課課綱主題所本的原書出版頁(ISBN 9780262048644,2023-12 出版)
- udlbook 官方網站(作者釋出的 PDF、投影片與習題) — 原書作者維護的免費資源站(udlbook.com 會轉址到此)