L4.1 生產線出了瑕疵,怎麼知道是哪一站的責任?
本篇重用M0.3漂流的溫度計:全導數、微分穿過積分與 JVP·M0.0霧中下山:Gradient 與方向
起點:一條長長的生產線
一條生產線有 50 站。最後出來的成品有瑕疵,扣了 3 分。
問題不是「誰有錯」,而是「每一站各要負多少責任」——更精確地說:第 17 站的溫度如果調高 0.1 度,最後的扣分會變多還是變少、變多少?
有一個笨辦法:把第 17 站調高一點,整條線重跑一次,看扣分差多少。 這行得通,但 50 站就要重跑 50 次;如果每一站有 20 個旋鈕,就是 1000 次。
而這正是我們的處境。霧中下山,一步該走多大:梯度下降的一頁 以來每一個方法都需要 ,而 有幾萬、幾百萬、幾十億個分量。逐個試一遍是不可能的。
一條長長的生產線出了瑕疵,
怎麼一次算出每一站的責任,
而不是每一站各重跑一次?
課堂提問Q1
把「每一站的責任」翻成數學。鏈式法則給了一串連乘——那為什麼還會有「便宜」與「貴」的差別?兩種乘法順序各是什麼,各要多少運算?
先想一想,再展開看整理後的答案
先寫下來。 一個 層的網路是一串函數的合成:
鏈式法則(漂流的溫度計:全導數、微分穿過積分與 JVP)給
式子本身沒有歧義,但「先乘哪一對」有。 矩陣連乘滿足結合律,所以答案一樣;成本不一樣。
從左往右(反向模式)。 先算最左邊那個 的列向量,再一路往右乘:
每一步的成本是 (一個長度 的向量乘一個 的矩陣),走完 層是 ——和一次前向傳播同一個量級。而且走這一趟的過程中,每一層的 都被算出來了,所以所有層的參數梯度一次拿齊。
從右往左(前向模式)。 先把中間那串 Jacobian 乘起來,每一步是 乘 ,成本 ,走完是 ——多一個 倍。
關鍵是損失是純量。 最左邊那一項是 的列向量,右邊則是 的大矩陣。從窄的那一端開始乘,中間產物一直是向量;從寬的那一端開始,中間產物是矩陣。
用生產線的話說:從成品往回問「這一站多影響一點會怎樣」,一次就把整條線的責任分完;從每一站往前推「我動一下會怎樣」,每一站都要推一遍。
這也解釋了為什麼「 個參數要 次」的直覺是錯的——那個直覺對應的是前向模式(每次推一個參數),而反向模式把方向倒過來之後, 消失在成本裡了。
兩個方向,兩種問句
自動微分的兩種模式,可以用它們各自回答的問句記住:
- 前向模式(JVP,Jacobian-vector product):「輸入沿著這個方向動一點,輸出會怎麼動?」一次前向掃描,成本 一次前向,回答一個輸入方向。輸入有 個維度就要掃 次。
- 反向模式(VJP,vector-Jacobian product):「輸出的這個組合,對每一個輸入各有多敏感?」一次反向掃描,成本 常數倍的前向,回答一個輸出方向,但一次給出全部 個輸入的答案。
於是規則很簡單:
訓練時 (損失是一個純量)而 是幾百萬,所以反向壓倒性地便宜。但這個結論會隨著問題翻轉:如果你要的是一整個 Jacobian( 很大),或者輸入只有兩三個維度,前向模式反而便宜。
「反向傳播」就是「把反向模式自動微分用在 的損失上」,沒有多一分神秘。
一次反向就把全部拿到
把上面的成本算式驗證一次。同一個網路、同一個位置、同一個梯度,兩種算法:
# 反向傳播:一次前向 + 一次反向
G, _ = mlp.grads(P, X, y)
# 有限差分:每個參數各擾動兩次,重跑整個前向
for i in range(n_params):
vp = v0.copy(); vp[i] += eps
vm = v0.copy(); vm[i] -= eps
g_fd[i] = (loss(unflat(P, vp)) - loss(unflat(P, vm))) / (2*eps)
參數 33 反向傳播 0.38 ms 有限差分 5.2 ms 倍數 13.8× 相對差 2.22e-10
參數 321 反向傳播 0.26 ms 有限差分 32.5 ms 倍數 123.0× 相對差 4.29e-10
參數 2497 反向傳播 0.35 ms 有限差分 294.3 ms 倍數 840.5× 相對差 8.07e-10
三件事同時被證實了。
反向傳播的時間幾乎不動( ms——差異是量測噪聲,不是趨勢),而參數從 33 變成 2497。這正是「成本 一次前向」那句話:前向本身也隨參數變多而變慢,但它是同一次掃描,不是 次。
有限差分線性成長( ms,參數約 /,時間約 /),倍數因此一路長到 。真實網路的 是 到 ,那個倍數就是同樣的數量級——不是「慢一點」,是「不可能」。
兩者算出來的是同一個梯度(相對差 )。這一欄很重要:它說明反向傳播不是近似,是精確的導數(在浮點誤差之內)。它和有限差分的差別純粹是算法,不是精度取捨。
反向傳播不是一個近似技巧,是同一個導數的一種便宜算法。
補充代價是記憶體:activation 要留著
反向那一趟需要每一層的中間值(前向算出來的 ),因為 Jacobian 依賴它們。所以反向傳播用時間換記憶體:時間是一次前向的常數倍,記憶體是 。
這在大模型上是真正的瓶頸,於是有一整套對策:gradient checkpointing(只存一部分 activation,反向時把缺的重算一次——時間多約 ,記憶體降一個數量級)、混合精度(activation 存成半精度)、以及重算比存還便宜的 layer 設計。
值得記住的形狀是:反向傳播把「 次前向」換成了「一次前向 + 一次反向 + 一份 activation」。前兩項是時間,第三項是記憶體——而深度學習框架的工程史,有一大半是在管理第三項。
什麼時候還是要用數值梯度
反向傳播全面勝出之後,有限差分只剩一個用途,但那個用途很重要:檢查你自己寫的反向對不對(gradient checking)。手寫一層新的運算時,反向的公式很容易差一個轉置、一個因子或一個和的軸。
做法是挑幾十個座標,比對兩者。而這時步長 要選對——它有一個 U 形:
中央差分的相對誤差對步長(48 寬的網路,和反向傳播比)
h=1e-01 1.087e-03
h=1e-03 1.091e-07
h=1e-05 1.671e-11 ← 最好
h=1e-07 1.449e-09
h=1e-09 1.343e-07
h=1e-11 1.223e-05
兩邊各有一個誤差來源,方向相反。
左邊是截斷誤差。 中央差分的 Taylor 展開(導航的三十秒:Taylor 展開與「假設直線」 的三階項)給
誤差 。實測從 到 誤差掉了約 倍()——正好是 。
右邊是捨入誤差。 分子是兩個接近的數相減,雙精度浮點只有約 16 位有效數字,所以分子的絕對誤差大約是 ,除以 之後是 。實測從 到 誤差長了約 倍——正好是 。
兩條線在 交會,那裡的誤差是 。這也給出一條實務規則:gradient checking 用 (雙精度),並且看相對誤差;小於 算通過,大於 幾乎一定是公式錯了。
有限差分的精度有一個下限,而且那個下限不在 。
注意用單精度做 gradient checking 會失敗
上面的 U 形谷底位置由浮點的有效位數決定。雙精度(float64)約 16 位,谷底在 、誤差 ;單精度(float32)只有約 7 位,谷底移到 附近而且谷底的誤差只到 量級。
後果是:在單精度下做 gradient checking,正確的實作也會「失敗」——你看到 的相對誤差,分不出那是捨入還是真的寫錯。所以規則是:gradient checking 一律在雙精度下做,而且只檢查那一層的運算,不要連著整個大模型一起檢查(層數越多, 的量級越大,捨入誤差跟著放大)。
回到情境:每一站的責任怎麼被分下去
回到那條生產線。反向那一趟在做的事,逐站來看是:
- 最後一站:成品扣了 3 分, 是「輸出每動一單位,扣分動多少」。
- 往前一站:拿到後面傳來的「你的輸出有多重要」(一個向量),乘上這一站的 Jacobian,得到「我的輸入有多重要」——傳給更前面一站。
- 同時:把那個向量乘上 ,得到這一站自己的旋鈕該怎麼調。
每一站只需要知道兩件事:自己的局部導數,以及後面傳回來的那個向量。它完全不需要知道整條線長什麼樣子。這就是為什麼深度學習框架可以讓你任意接一張計算圖——每個運算只要實作自己的前向與 VJP,組合的事情框架會處理。
這一切依賴什麼。 一,每一站可微(實務上「幾乎處處可微」就夠了,ReLU 在 0 的那一點框架直接指定一個值)。二,計算圖是有向無環的——有迴圈就要展開(這正是 RNN 的 backpropagation through time,展開之後記憶體隨序列長度成長)。三,中間值留得住;留不住就要重算(見上面的 Remark)。
回到情境:mlp.py 裡那三十行
同一張地圖從不同地方出發:損失曲面與初始化 用的那個 numpy MLP,反向那一段現在可以逐行讀了:
def grads(P, X, y, skip=False):
out, hs = fwd(P, X, skip) # 前向:順便把每一層的中間值 hs 留著
g = 2*(out - y[:,None])/len(y) # ∂L/∂h_D:損失對最後一層輸出的導數
G = []
for i in range(len(P)-1, -1, -1): # 從最後一層往前走
W, b = P[i]; h = hs[i]
G.append((h.T @ g, g.sum(0))) # 這一站自己的責任:∂L/∂W_i、∂L/∂b_i
if i > 0:
d = g @ W.T # 往前傳:∂L/∂h_{i-1} 的線性部分
z = hs[i-1] @ P[i-1][0] + P[i-1][1]
d = d * (1 - np.tanh(z)**2) # 再乘上 tanh 的導數
g = d
return list(reversed(G)), ((out - y[:,None])**2).mean()
三件事對得上前面的推導。第一行把 hs 留下來——就是 Remark 講的那份記憶體代價。g 從頭到尾是一個向量(每一筆資料一列),不是矩陣——這就是「從純量那一端開始乘」。迴圈裡每一圈做兩件事:把這一站的梯度收起來(G.append),把責任往前傳(g = d)——正好是上一節生產線的第 2 與第 3 點。
h.T @ g 這一行值得停一下:它同時完成了「對 batch 求和」與「 的乘法」。因為 ,所以 ——前向用到的那個中間值,在反向時變成乘數。這就是為什麼 activation 非留不可。
刻意違反:把那一行 拿掉。 這是手寫反向最常見的錯誤之一(忘記乘激發函數的導數)。這種 bug 不會讓程式崩潰,也不會讓損失完全不降——它給出的是一個「方向大致對但比例錯掉」的向量,訓練看起來還在動,只是慢很多、而且對學習率異常敏感。
這正是 gradient checking 存在的理由:上面的 U 形告訴你,用 比對幾十個座標,正確的實作會給 量級的相對誤差,漏乘一項會給 量級——兩者差十個數量級,不可能看錯。
兩個方向的掃描,與有限差分的 U 形
互動 demo:反向傳播那條線幾乎水平,有限差分那條線性往上;下面的 V 形谷底大約在 h = 10⁻⁵。
先消化一下
參考文獻
- Rumelhart, D., Hinton, G., Williams, R. Learning representations by back-propagating errors. Nature 1986.(把反向傳播帶進神經網路的那篇;它的貢獻不是鏈式法則,是把它組織成一個可以套在任意深度上的演算法。)
- Griewank, A., Walther, A. Evaluating Derivatives. 2nd ed. SIAM 2008.(自動微分的標準參考:前向/反向模式的成本分析、checkpointing 的理論。)
- Baydin, A. et al. Automatic Differentiation in Machine Learning: a Survey. JMLR 2018.(把 AD 與數值微分、符號微分分清楚;本篇「反向傳播不是近似」那一節的來源。)