自動微分(automatic differentiation)
在電腦上求導有三條路,而它們容易混淆。符號微分操弄公式(並膨脹成龐大的表達式)。有限差分用小步長 h 戳函數(並受步長取捨之苦,永遠不全然精確)。自動微分是第三條、現代的路:它把鏈鎖律機械地套用在你程式實際執行的運算序列上,「精確地」算出導數——達到完整的機器精度——沒有步長,也不操弄公式。
關鍵洞見是:程式所算的任何函數,無論多複雜,都由一張基本運算(加、乘、sin、exp ...)的圖組成,而這些運算各自的導數是已知的。自動微分藉鏈鎖律讓導數資訊在這張計算圖中傳播。「前向」模式在每個值旁邊攜帶它對某個輸入的導數(運作上,把程式跑在對偶數 a + b*epsilon 上,其中 epsilon^2 = 0,epsilon 部分浮現為導數);當輸入少、輸出多時它有效率。「反向」模式先把程式前向跑一遍記下圖,再反向掃一遍,一次累積出某個輸出對每個輸入的導數;當輸入多、輸出一個時它有效率。那個反向模式正是反向傳播——訓練神經網路的引擎,那裡你要的是單一純量損失對數百萬個參數的梯度。
自動微分是有限差分的現代、精確替代,也是基於梯度的最佳化與深度學習能擴展的原因:無論參數多少,單次反向掃描就以約一次函數計算的成本給出整個梯度。誠實的細節:它精確地微分程式碼「所算的東西」,包括那段程式碼裡的任何錯誤或近似;它在不可微之處需要小心(abs(x) 的折角或 relu 在 0 處、分支、長度依資料而定的迴圈);反向模式必須儲存中間值,所以在很深的計算上記憶體可能成為瓶頸;而且它給的是導數,不是函數的積分——它是本領域中與求積一半相對的微分一半。
用前向模式的對偶數對 f(x) = x * sin(x) 在 x = 2 求導。把 x 帶成 (2, 1),意指值 2、導數 1。則 sin 給出 (sin 2, cos 2 * 1) = (0.909, -0.416),而 (2,1)*(0.909,-0.416) 上的乘法律得出值 1.819、導數 2*(-0.416) + 1*0.909 = -0.926——正是精確的 f'(2),完全沒有步長。
讓鏈鎖律走過計算圖:精確,沒有步長 h。
自動微分「不是」有限差分,也「不是」符號代數——它沒有步長誤差,也沒有表達式膨脹。但它微分的是寫成的程式碼:在不可微之處(abs、relu 在 0、一個 if 分支)它回傳某一個次梯度,而非警告;而反向模式儲存中間值的記憶體,在很深的圖上可能成為主宰。