訓練分類模型,最後總要把模型的輸出變成一個能調整參數的訊號,而這一步幾乎清一色都用交叉熵(cross-entropy)。回歸問題可以直接拿預測值跟目標值相減、平方,分類很少這麼算。會這樣選,是因為 softmax 接上交叉熵之後,對模型輸出求導會得到一個特別乾淨的結果。下面就把導數算出來,順便看它為什麼會長成這個形狀。
先擺一個小例子在桌上。要分貓、雞、狗三類,手上的圖其實是貓,正確答案寫成 one-hot 向量 y = (1, 0, 0)。模型看完圖,對三類各給一個分數,這三個分數就是 logit,習慣寫成 o = (o₁, o₂, o₃),比如 o = (2.0, 1.0, 0.1)。logit 本身不是機率,可以是負的,三個加起來也不保證等於 1。
為什麼不乾脆把三個分數當成預測,直接跟 (1, 0, 0) 相減再平方,像回歸一樣?Dive into Deep Learning 的 softmax 回歸章節講得很白:把分類當成向量對向量的回歸問題,其實「surprisingly well」,但不理想。問題出在 logit 不具備機率該有的性質,既不保證非負,加起來也不保證是 1。硬拿它們做平方誤差,等於用沒有範圍的東西去逼近機率向量。所以中間得先加一道手續,把 logit 壓成一組合法的機率,這道手續就是 softmax:每個分數取指數,再除以全部指數的總和。套進上面的數字,o = (2.0, 1.0, 0.1) 過完 softmax 大約是 (0.66, 0.24, 0.10),三個都落在 0 到 1 之間,加起來剛好是 1。
有了機率之後,交叉熵的定義其實很短:把每一類的預測機率取對數、乘上該類的標籤,全部加起來再取負號,寫成 l = −Σₖ yₖ log ŷₖ。標籤是 one-hot 的時候,y 只有正確類別是 1、其餘是 0,加總裡其他項都被乘成 0,所以整條損失只剩下一項:正確類別的機率取負對數。以這張貓的圖來說,模型給貓的機率是 0.66,損失就是 −log(0.66),大約 0.42。假如模型很有把握、給貓 0.99,損失掉到 −log(0.99) 約 0.01;假如模型幾乎不信是貓、只給 0.01,損失衝到 −log(0.01) 約 4.6。交叉熵只盯著正確類別的機率,機率越低、罰得越重,其他兩類給多少並不直接管。
真正值得停下來的是下一步:把 softmax 接上交叉熵,然後對 logit 求導。推導之前要先把完整的加總寫回來。剛才算數字時走了 one-hot 的捷徑,只留正確類別一項;推導得從 l = −Σₖ yₖ log ŷₖ 出發,每一類都留著,才看得出化簡後的式子是怎麼來的。把 softmax 的定義代進去,加總裡每一項成了 −yₖ log(exp(oₖ) / Σᵢ exp(oᵢ));分母那個加總跑遍所有類別,跟外層是兩件事,所以換個代號寫成 i。這一步靠的是對數的除法規則,log(a/b) 等於 log a − log b,前面又帶著負號,於是 −log(exp(oₖ) / Σᵢ exp(oᵢ)) 就是 log Σᵢ exp(oᵢ) − oₖ,其中 log exp(oₖ) 直接還原成 oₖ。代回加總,損失拆成兩塊:Σₖ yₖ log Σᵢ exp(oᵢ) − Σₖ yₖ oₖ。前一塊裡的 log Σᵢ exp(oᵢ) 跟 k 無關,每一項乘的都是同一個數,可以提到加總外面,留下的 Σₖ yₖ 正好是 one-hot 向量各項相加,等於 1。整條損失於是化簡成 l = log Σᵢ exp(oᵢ) − Σₖ yₖ oₖ。
接下來挑其中一個 logit oⱼ 求偏導。後一塊 Σₖ yₖ oₖ 攤開來是 y₁o₁ + y₂o₂ + y₃o₃,其中只有 yⱼoⱼ 含有 oⱼ,其餘各項對 oⱼ 而言都是常數、導數為 0,所以整塊的偏導就是 yⱼ。前一塊 log Σᵢ exp(oᵢ) 要用連鎖律:外層 log u 的導數是 1/u,u 就是整個指數和 Σᵢ exp(oᵢ);內層 Σᵢ exp(oᵢ) 對 oⱼ 求導,同樣只有 exp(oⱼ) 含有 oⱼ,其餘都是常數,所以內層留下 exp(oⱼ)。兩層相乘是 exp(oⱼ) / Σᵢ exp(oᵢ),也就是 softmax 給第 j 類的機率。前後兩塊相減,∂l/∂oⱼ = softmax(o)ⱼ − yⱼ。整段推導沒用到什麼技巧,log 和 exp 在中間互相抵消掉大半,剩下的就是這條式子。
停在這個結果上。梯度白話講就是「模型給某一類的機率,減掉該類到底有沒有發生」。回到貓的圖,y = (1, 0, 0)、softmax 是 (0.66, 0.24, 0.10),三個 logit 收到的梯度就是 (0.66−1, 0.24−0, 0.10−0),即 (−0.34, 0.24, 0.10)。正確類別(貓)收到負梯度,梯度下降會把貓的 logit 往上推;另外兩類收到正梯度,會被往下壓。d2l 說這個形狀跟線性回歸是同一個:回歸算的是觀測值減預測值,分類算的是預測機率減實際發生,兩邊長得像並不是巧合。對實作來說,最實際的好處就是梯度好算。
再多看一眼式子,會發現它自己就把「錯得多離譜」量化好了。梯度既然是機率減標籤,大小就天生有界,也好讀。一個根本沒發生的類別,y = 0,模型不管怎麼亂給,梯度頂多接近 1;模型越是自信地把機率堆在錯的類別上,往回修的力道就越接近上限。反過來,正確類別已經給到接近 1 的時候,softmax(o) − y 接近 0,幾乎不再更新。信心錯得越離譜、梯度就越大,這是 softmax(o) − y 直接算出來的結果,沒有另外加規則去調整力道。
實作上不用自己動手串 softmax。PyTorch 的 nn.CrossEntropyLoss 官方文件寫的是,它吃的是 logit,「the unnormalized logits for each class」;criterion 內部等價於先做 LogSoftmax 再接 NLLLoss,softmax 那一步已經包在裡面了。所以模型最後一層直接輸出原始分數就好:
import torch
import torch.nn as nn
loss_fn = nn.CrossEntropyLoss()
# 3 筆資料、5 個類別;logits,沒有先過 softmax
logits = torch.randn(3, 5, requires_grad=True)
target = torch.tensor([1, 0, 4]) # 每筆是類別索引
loss = loss_fn(logits, target)
loss.backward()
先自己套一層 softmax 再丟進 CrossEntropyLoss,softmax 就做了兩次,梯度也不再是上面推出來的乾淨式子。
回到開頭那張貓的圖。假如模型不溫吞地給 0.66,而是非常有把握地認定圖裡是雞,機率壓到 (0.01, 0.98, 0.01),三個 logit 拿到的梯度就是 (0.01−1, 0.98−0, 0.01−0),也就是 (−0.99, 0.98, 0.01)。貓的 logit 被接近 1 的力道往上抬,雞的 logit 被幾乎同樣的力道壓下去。修正的方向完全由 y 決定。如果 (1, 0, 0) 本身標錯了,這條乾淨的訊號會用同樣大的力道把模型往錯的方向拉,公式並不知道標籤是不是真的。



