
今天是我在浙大跟著學習小組推進機器學習的第18天。前17天從線性回歸、邏輯回歸一路做到簡單的MLP二分類已經玩得比較順手了所以我一度以為多分類問題不過就是換個標簽、加幾個輸出節點的事。結果真正動手做才發現多分類不只是把輸出維度從1改成K這么簡單它牽扯到輸出層的激活方式、損失函數的設計、評估指標的選取還有一堆訓練時才暴露出來的隱藏坑點。這篇文章就把我Day18這天在疏錦行項目里學習和實戰多分類問題的完整過程記錄下來。內容包括多分類與二分類的本質差異、Softmax與交叉熵的配合邏輯、評估指標的選取原則、一個完整的PyTorch多分類實戰流程以及我當天踩過的幾個比較典型的坑。如果你正在學習機器學習分類任務或者從二分類往多分類邁進時覺得哪里差一口氣這篇文章應該能幫你少走不少彎路。1. 為什么多分類值得單獨花一天二分類順手多分類翻車很多教程把邏輯回歸講完之后緊接著就說一句多分類就是把邏輯回歸擴展成Softmax回歸然后這頁PPT就翻過去了。我以前也是這么覺得的直到自己動手訓練一個10分類模型才發現其中的細節比想象中多得多。這一節先不聊代碼我把多分類問題的認知框架梳理清楚。1.1 二分類到多分類模型輸出結構發生了什么變化二分類問題里面模型輸出的是一個標量或者一個二維概率向量。以邏輯回歸為例輸出經過Sigmoid之后落在0到1之間這個值可以理解為屬于正類的概率負類的概率就是1減去這個值。但到了多分類場景假設有K個類別模型的輸出就不再是一個數而是一個K維向量。絕大多數情況下我們要求這K個維度上的值加起來等于1并且每一維的值都在0到1之間這樣它才能被解釋成模型對這個樣本屬于第k個類別的置信概率。如果你只是簡單地把輸出維度改成K然后仍然用Sigmoid對每個維度單獨做激活這樣得到的每一維雖然都在0到1之間但它們的總和并不等于1。這個結果其實是K個獨立的二分類概率不是真正意義上的多分類概率分布。Day18這天我最開始就犯了這個錯誤后面在損失計算時怎么調都覺得數值怪怪的。這也是為什么多分類必須有專門的Softmax層或者類似機制來保證概率歸一化。1.2 多分類的多種落地路徑OvR、OvO與Softmax回歸多分類可以走的路不止一條。最容易理解的是拆解法把K分類問題拆成若干個二分類問題。一是一對多OvROne-vs-Rest訓練K個二分類器第k個分類器負責區分屬于第k類和不屬于第k類預測時把所有分類器跑一遍取置信度最高的那個作為最終類別。二是一對一OvOOne-vs-One對任意兩類都訓練一個分類器總共訓練K(K-1)/2個預測時用投票法決定類別這是SVM在處理多分類時常用的手段。第三種就是Softmax回歸也就是Logistic回歸的直接推廣。它不拆問題而是直接用一個模型輸出K維概率分布讓這K個概率之間互相競爭、彼此約束。神經網絡處理多分類問題時幾乎都走這條路線因為它端到端訓練、梯度和模型結構都更自然。我在做疏錦行項目時選擇的是Softmax回歸方案。坦白說OvR和OvO更容易理解但放在深度學習框架里Softmax的寫法最簡潔、訓練效率也最高而且后續不管是加正則化還是換更復雜的網絡結構都是在這一套框架上擴展。1.3 多分類問題的難度來源類別間的邊界競爭多分類比二分類難本質上是因為類別之間的邊界變復雜了。二分類只需要找一條決策邊界把空間切成兩半而K分類需要找到K個區域每個區域對應一個類別區域與區域之間的邊界可能是多條甚至可能是非線性的。更麻煩的是真實數據里面類別之間往往存在相似區域。比如CIFAR-10里面鳥和飛機都有翅膀貓和狗都是四條腿加毛茸茸這些類別特征重疊的地方就是模型最容易犯錯的地方。二分類任務里通常只關心是與不是而多分類任務里的錯誤是有遠近親疏關系的——把貓認成狗和把貓認成卡車雖然都是錯但前者在語義上顯然更接近。這種錯誤結構在評估模型時也需要額外關心。理解了這一點你就會明白為什么多分類要單獨做評估、單獨調損失函數你不能只看正確率這一個數字你需要知道模型在哪些類別上犯糊涂它們是在把哪個類別誤認成哪個類別。2. Softmax與交叉熵多分類模型的兩個核心齒輪多分類深度學習模型的標準配置是最后一層輸出K維向量 Softmax歸一化 交叉熵損失。這套組合拳不是憑空來的每一步都有它存在的理由。Day18我把這兩個東西的數學原理和代碼實現都過了一遍下面是我認為最關鍵的幾塊拼圖。2.1 Softmax的數學本質把得分變成概率分布假設模型最后一層全連接層輸出的原始分數是 z [z_1, z_2, ..., z_K]Softmax做的事情就是P(yi|x) exp(z_i) / Σ_{j1}^{K} exp(z_j)也就是說對每個分數取指數再除以所有指數之和。這樣做的效果有兩個第一所有輸出都是正數第二所有輸出加起來等于1滿足概率分布的定義。為什么要用指數而不是直接用 z_i 除以 z 的和因為原始分數可能是負的直接歸一化會出問題。而且指數運算會放大分數之間的差異讓原本得分最高的那個類別的概率更顯著這符合我們分類時的直覺得分稍微高一點就應該有更大概率被選中。但Softmax還有一個副作用值得注意它會把差距拉得特別大。假設三個類的得分是 [2, 1, 0]經過Softmax之后概率大約是 [0.665, 0.245, 0.090]。如果得分變成 [2, 1, 0.1]看起來差別不大但概率已經變成 [0.659, 0.242, 0.099]說明Softmax對小數點后的變化也比較敏感。這個特性在模型訓練后期會體現為模型的預測概率越來越高哪怕它其實沒那么確定。2.2 數值穩定性問題exp一不小心就溢出Softmax里面的 exp(z_i) 在 z_i 比較大的時候會爆炸。比如 z_i 1000exp(1000) 直接就是無窮大程序里會變成NaN。這在小模型里可能不常見但一旦沒做歸一化就直接輸入網絡經常見。解決辦法很簡單先找出 z 里面的最大值 m然后算 exp(z_i - m)。因為 Softmax 的分子分母同時除以 exp(m)結果不變但數值范圍被控制住了。這就是代碼里常見的那行z z - torch.max(z, dim-1, keepdimTrue).values p torch.exp(z) / torch.sum(torch.exp(z), dim-1, keepdimTrue)我在Day18的實踐里就把這一步寫進了自定義邏輯中。當然如果你直接用PyTorch的torch.nn.CrossEntropyLoss框架內部已經處理好了數值穩定性不用自己操心。但理解這個處理方式仍然很重要因為當你去讀別人的代碼、或者自己寫損失函數時這點小細節往往決定了訓練能不能穩定跑下去。2.3 交叉熵損失為什么它是多分類的首選有了概率分布之后我們需要一個損失函數來衡量模型給出的概率分布與真實標簽之間的差距。交叉熵的公式是L -Σ_{k1}^{K} y_k * log(p_k)其中 y 是真實標簽的獨熱編碼p 是模型預測的概率。因為 y 是獨熱編碼只有真實類別那一維是1其他都是0所以公式可以簡化為L -log(p_c)其中 c 是樣本的真實類別。直觀理解就是模型給真實類別分配的概率越高loss越小給真實類別分配的概率越低loss越大。這里有一個很關鍵的梯度性質Softmax之后接交叉熵它的梯度是?L/?z_i p_i - y_i這個公式非常漂亮。它意味著梯度的計算就是預測概率減去真實標簽不需要鏈式法則一層層地算。如果模型覺得某個類別的概率是0.8而真實類別其實是那個類別梯度就是負的0.2推動模型往降低該類別得分的方向更新。這個性質讓Softmax交叉熵在數值上非常穩定也是它們在分類任務里成為黃金組合的根本原因。2.4 在實際代碼里CrossEntropyLoss替你做了什么PyTorch的torch.nn.CrossEntropyLoss是一個封裝了三個功能的類LogSoftmax、負對數似然損失NLLLoss、以及一些內部優化。也就是說你在調用它的時候不需要在最后一層額外加Softmax直接把網絡輸出的logits丟進去就行。如果你在最后一層手動加了Softmax再把結果傳給CrossEntropyLoss等于算了兩次Softmax數值會變差訓練也可能出問題。我自己的教訓是網絡最后的輸出層應該保持裸的K維向量訓練階段用CrossEntropyLoss計算loss只有到了推理階段需要看概率分布時才在模型輸出后手動包一層Softmax。這個習慣從那之后一直沒變過。3. 多分類評估指標準確率之外我更該看什么多分類任務的評估是最容易被忽視的環節。初學者一般只看一個數字——Accuracy準確率——只要模型在測試集上達到90%就覺得萬事大吉。但Day18做完CIFAR-10那個例子之后我發現單單看準確率會漏掉很多信息尤其是當類別分布不均勻或者模型在特定類別上有系統性問題的時候。3.1 混淆矩陣看清每一類到底被認成了什么混淆矩陣是多分類評估的第一站。它是一個K×K的矩陣第 i 行第 j 列的含義是真實類別為 i、但被模型預測為 j的樣本數量。對角線上的數字越大越好非對角線上的數字就是具體的錯誤模式。CIFAR-10上訓練完成后我打印出混淆矩陣發現模型很容易把狗預測成貓把鹿預測成馬。這些錯誤本身有高度的結構相似性——四足動物之間互相混淆。這種信息在準確率數字里完全體現不出來但它恰恰指導著我們下一步改進的方向是需要給某些類別加更多訓練樣本還是需要設計更好的特征提取器來區分相似類別。在代碼層面可以用sklearn.metrics.confusion_matrix一行計算也可以用PyTorch在測試循環里自己累加confusion torch.zeros(num_classes, num_classes, dtypetorch.long) for x, y in test_loader: pred model(x).argmax(dim1) for t, p in zip(y.view(-1), pred.view(-1)): confusion[t, p] 1有了混淆矩陣你可以非常直觀地定位病根。3.2 Macro-F1、Micro-F1和Weighted-F1到底該選哪個由于準確率在多類不均衡時缺乏參考價值更合理的做法是看F1。但F1在多分類環境里有三套計算方式我第一次看的時候也容易暈這里用大白話梳理一下。Macro-F1宏平均先對每個類別單獨計算精確率和召回率得到一個F1然后把所有類別的F1取平均。它對每個類一視同仁不會因為某個類樣本多就占更大權重。如果你特別關心小類別能不能被識別出來Macro-F1是更嚴格的指標。Micro-F1微平均把所有類別的預測結果匯總到一起全局統計TP、FP、FN然后計算整體的精確率和召回率最后算F1。當所有類別的樣本量差不多時Micro-F1和Accuracy數值上會非常接近當類別不均衡時大類別會主導Micro-F1。Weighted-F1加權平均還是先對每個類算F1但按每個類的真實樣本占比給它加權然后求和。這樣既保留了逐類的F1信息又反映了類別在數據中的實際重要性是把Macro和Micro各取一半的做法。樣本不均衡時我優先推薦它。3.3 不均衡多分類場景的處理思路現實中很多多分類任務不同類別的樣本數量差異很大比如故障診斷里正常樣本遠多于故障樣本圖像識別里某些罕見物種幾乎沒啥訓練數據。這時候直接訓練出來的模型會傾向于把所有樣本都預測成大類準確率看著挺高實際一點用都沒有。處理思路通常有幾種第一種是對損失函數加權給樣本少的類別更大的權重第二種是過采樣復制小類別的樣本讓它們多出現幾次第三種是改用Focal Loss把焦點放在那些難以分類的樣本上第四種是在評估時堅持看Macro-F1而不是Accuracy。Day18我的CIFAR-10數據還算均衡所以沒有走這些復雜路子但我把Focal Loss的原理搞清楚了就是給交叉熵加一個調制因子 (1-p_t)^γ當模型已經能把某個樣本分得很好時讓這個樣本產生的梯度權重降低迫使模型更多關注那些老大難樣本。后續如果遇到不均衡數據這套思路可以直接平移過去。4. PyTorch多分類實戰用CIFAR-10把概念跑通理論說得再多不如一個完整案例直接。Day18下午我用PyTorch搭了一個CIFAR-10圖像分類模型從數據準備到訓練再到測試把多分類的完整鏈路跑了一遍。下面把步驟和關鍵代碼貼出來附帶每一步的解釋方便你照著復現。4.1 數據準備CIFAR-10與數據增強的取舍CIFAR-10是一個10類、每類6000張32×32彩色圖像的經典數據集類別包括飛機、汽車、鳥、貓、鹿、狗、青蛙、馬、船、卡車。它規模不大訓練集5萬張、測試集1萬張非常適合在學習階段把多分類的流程跑通。import torch import torchvision import torchvision.transforms as transforms transform_train transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) transform_test transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)), ]) trainset torchvision.datasets.CIFAR10(root./data, trainTrue, downloadTrue, transformtransform_train) testset torchvision.datasets.CIFAR10(root./data, trainFalse, downloadTrue, transformtransform_test) trainloader torch.utils.data.DataLoader(trainset, batch_size64, shuffleTrue, num_workers2) testloader torch.utils.data.DataLoader(testset, batch_size64, shuffleFalse, num_workers2)這里有個細節值得多說一句數據增強只加在訓練集測試集只用ToTensor和Normalize。RandomCrop和RandomHorizontalFlip是在訓練時給模型變著花樣看數據增強泛化能力但測試時必須保證數據的真實性和一致性否則測試結果會失真。Normalize那三個數分別是CIFAR-10數據集在RGB三個通道的均值和標準差。歸一化不是可有可無的體操動作它能讓輸入特征的數值范圍穩定在0附近幫助模型更快收斂也能避免某些特征數值過大導致梯度不穩定。4.2 網絡結構一個足夠應付CIFAR-10的簡單CNN我沒有一開始就上ResNet而是先用一個三層卷積加全連接的小網絡把多分類流程跑通再說。網絡結構是這樣的import torch.nn as nn import torch.nn.functional as F class SimpleCNN(nn.Module): def __init__(self, num_classes10): super(SimpleCNN, self).__init__() self.conv1 nn.Conv2d(3, 32, kernel_size3, padding1) self.conv2 nn.Conv2d(32, 64, kernel_size3, padding1) self.conv3 nn.Conv2d(64, 128, kernel_size3, padding1) self.pool nn.MaxPool2d(2, 2) self.fc1 nn.Linear(128 * 4 * 4, 256) self.fc2 nn.Linear(256, num_classes) self.dropout nn.Dropout(0.3) def forward(self, x): x self.pool(F.relu(self.conv1(x))) x self.pool(F.relu(self.conv2(x))) x self.pool(F.relu(self.conv3(x))) x x.view(x.size(0), -1) x F.relu(self.fc1(x)) x self.dropout(x) x self.fc2(x) return x32×32的輸入經過三次2倍池化之后空間尺寸變成4×4所以全連接層第一層的輸入維度是128×4×42048。fc2的輸出是10維對應CIFAR-10的10個類別。注意forward的最后沒有Softmax訓練時直接把原始logits丟給CrossEntropyLoss。在動手寫網絡前算一算特征圖尺寸是個好習慣這樣能減少調試維度不匹配的時間。如果懶得算也可以先在代碼里打印一句print(x.shape)確認一下。4.3 訓練循環與優化器選擇優化器我選了Adam學習率設成0.001這是很多小模型的穩妥起點。如果你選的SGD往往需要配合動量并且手動調學習率Adam對這種快速驗證的場景更友好。import torch.optim as optim model SimpleCNN(num_classes10) criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr0.001) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size10, gamma0.1) epochs 30 for epoch in range(epochs): model.train() running_loss 0.0 correct 0 total 0 for images, labels in trainloader: optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() scheduler.step() train_acc 100.0 * correct / total print(fEpoch {epoch1}/{epochs}, Loss: {running_loss/len(trainloader):.4f}, Acc: {train_acc:.2f}%) print(Training finished.)每訓練完一個epoch我還會順便在測試集上跑一遍記錄測試準確率方便觀察有沒有過擬合。StepLR讓學習率每10個epoch衰減10倍后期收斂更穩。有個小建議訓練循環里最好保留model.train()和model.eval()的顯式切換。因為Dropout和BatchNorm在訓練和推理時的行為不一樣忘了切會導致測試集上的指標偏低或波動。4.4 測試評估從準確率延伸到混淆矩陣和F1測試階段我把模型切到eval模式關閉梯度計算然后統計準確率、逐類精確率/召回率/F1并生成混淆矩陣。from sklearn.metrics import classification_report, confusion_matrix, f1_score model.eval() all_preds [] all_labels [] with torch.no_grad(): for images, labels in testloader: outputs model(images) _, predicted torch.max(outputs, 1) all_preds.extend(predicted.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) print(fTest Accuracy: {sum(1 for p, l in zip(all_preds, all_labels) if p l) / len(all_preds):.4f}) print(classification_report(all_labels, all_preds, digits4)) cm confusion_matrix(all_labels, all_preds) print(cm)classification_report會直接給你每一類的precision、recall、f1-score以及Macro和Weighted的平均值省去自己算的功夫。有了這些數字我就能判斷模型在哪些類別上弱再回頭去找原因。跑完30個epoch這個簡單CNN在測試集上大概能到78%左右的準確率。這跟當下動輒95%以上的大模型沒法比但作為理解多分類全流程的載體已經足夠了。真正有價值的不是數字而是從數據加載到評估指標這條鏈路的完整性和各個環節的處理邏輯。5. 多分類訓練中踩過的坑Day18實測的真實教訓這部分我特別想寫因為Day18訓練的多數時間其實不是在寫網絡而是在處理各種莫名其妙的報錯和反常現象。我把當天踩過的坑按癥狀—原因—解法的方式整理出來希望你能繞開。5.1 標簽格式獨熱編碼 vs 類別索引第一次寫多分類代碼時我習慣性地把標簽做成獨熱編碼再送進模型。然后發現CrossEntropyLoss報錯提示目標類別范圍不對。這里要特別強調PyTorch的CrossEntropyLoss要求的目標是整數索引也就是0到K-1之間的整數張量形狀一般是(B,)或(B, 1)而不是(B, K)的獨熱編碼。這是很多從Keras或其他框架轉過來的人常踩的坑——TensorFlow的categorical_crossentropy經常搭配獨熱編碼而PyTorch默認用整數索引。如果你確實已經生成了獨熱編碼需要轉回整數索引用torch.argmax(y_onehot, dim1)即可。如果你更習慣用獨熱編碼的那種寫法也可以自己調用torch.nn.functional.binary_cross_entropy_with_logits單獨設計損失但那樣又回到多個獨立二分類的套路不是標準多分類了。5.2 Loss直接變成NaN學習率過高和數值溢出我在跑一個實驗時把學習率調到了0.01結果不到3個epoch損失值就一路狂飆變成NaN。原因很簡單梯度更新步長太大參數一下跳到了損失函數曲面非常陡峭的區域梯度進一步爆炸最后數值直接溢出。排查這類問題我的經驗是先從這幾個方向下手。第一步把學習率先降回0.001看看loss是否恢復正常第二步檢查輸入數據里是否有NaN或異常大值歸一化是否正確第三步確認最后一層輸出沒有手動加Softmax再送給CrossEntropyLoss第四步如果網絡非常深可以考慮加梯度裁剪比如torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)。Loss變成NaN通常不是你寫得不夠好而是哪個數學操作在數值上不夠穩。按順序排查不要瞎試。5.3 輸出層到底要不要加激活函數多分類網絡的最后一層最常見的寫法是不加任何激活函數直接返回logits。很多人初學時會覺得多分類需要概率輸出所以最后一層要Softmax于是把Softmax放進網絡forward函數的最后。如果模型只是用來做推理這沒問題但如果訓練時還繼續用CrossEntropyLoss就會出問題。因為剛才說過CrossEntropyLoss內部自帶LogSoftmax和NLLLoss你在外面已經做過一次Softmax等于先算了一次概率分布又被取了一次對數這個對數再被當作logits放到內層的LogSoftmax里數值全都亂了。我的建議是訓練用的model只輸出logits不輸出概率。推理時再單獨對模型輸出執行Softmax或者直接torch.argmax(outputs, dim1)取預測類別連Softmax都可以省——因為Softmax是單調的不會改變argmax的結果。大多數情況下你需要的只是預測類別而不是精確的概率值。5.4 類別不平衡模型變成復讀機我在一個小型自定義數據集上做過測試其中類別A占了80%類別B和C各占10%。模型訓練完后準確率高達78%但一看混淆矩陣類別B和C幾乎全軍覆沒A類準確率95%以上。如果只看準確率你甚至會以為模型還不錯但它實際上只會無腦復讀A類完全失去了分類的意義。遇到類別不平衡的多分類優先做這幾件事計算每個類別的樣本數打印出來讓自己心里有數把loss的weight參數設置成更重視小類別的值評估時看Macro-F1而不是Accuracy如果數據量允許對小類別做過采樣或數據增強。PyTorch的CrossEntropyLoss自帶weight參數只需要傳入一個長度等于類別數的張量class_weights torch.tensor([0.8, 1.0, 2.0, 1.0, 1.5, 1.0, 2.0, 1.0, 1.0, 1.2]) criterion nn.CrossEntropyLoss(weightclass_weights)這樣一來錯分小類別的懲罰更大模型就有動力去學習小類別的特征了。5.5 每個類別的樣本數量、驗證集與隨機種子還有一個特別不起眼、但影響結果穩定性的點隨機種子。多分類模型的初始化權重、數據加載順序、數據增強的隨機性都會影響最終指標。如果你跑兩遍結果差很多第一件事就是固定隨機種子。Day18結束時我在訓練代碼開頭加了這幾行import random import numpy as np def set_seed(seed): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) set_seed(42)別小看這個動作它能讓你的實驗結果可復現在調參的時候避免被隨機性誤導。6. 從Day18往后看多分類進階的三個方向路跑通之后我并沒有急著繼續前進而是花了一點時間梳理多分類問題還能往哪些方向深入。這里把三個我認為最值得關注的方向列出來供同樣學到這里的朋友參考。6.1 Focal Loss與難樣本挖掘Day18用標準交叉熵跑CIFAR-10時模型最后卡在78%左右上不去很大一部分原因是那些容易混淆的樣本沒有被特別照顧。標準交叉熵對已經分類正確且高置信度的樣本也會產生梯度導致模型把大量精力浪費在簡單樣本上。Focal Loss做的事情就是壓低簡單樣本的貢獻權重讓模型集中火力處理困難樣本。它在交叉熵前面乘了一個調制因子FL -α_t * (1 - p_t)^γ * log(p_t)當 p_t 接近1時候系數接近0該樣本對loss幾乎沒有貢獻當 p_t 偏低時系數較大貢獻被保留。γ一般取2α_t用來調節正負樣本不平衡。這個損失函數最初出現在目標檢測領域但現在已經被廣泛用在各種不均衡分類任務中。6.2 標簽平滑讓模型別那么自信交叉熵損失為了最小化loss會逼著模型把真實類別對應的概率推向1。但在訓練數據本身存在噪聲或標注錯誤時這種過度自信會讓模型記憶噪聲反而降低泛化能力。標簽平滑的做法是把真實的獨熱編碼修改為y_smooth(k) 1 - ε當 k c 時y_smooth(k) ε / (K - 1)當 k ≠ c 時這里的 ε 是一個很小的超參數通常取0.1。它告訴模型真實類別不一定是絕對正確的其他類別也可能有一點概率等于給模型施加了正則化限制了輸出概率過于極端。在很多圖像分類競賽和業務模型里標簽平滑都能穩定提升泛化性能是我比較推薦嘗試的進階技巧。6.3 從多分類到多標簽Sigmoid與Binary Cross Entropy的過渡多分類的另一條進階路線是多標簽分類。多分類里一個樣本只能屬于一個類別但很多真實場景里一個樣本可能同時擁有多個屬性標簽比如一張圖片里面同時有貓和狗或者一篇文章同時涉及科技、經濟兩個主題。多標簽的做法和標準多分類有個關鍵差異輸出層不再用Softmax歸一化而是用Sigmoid對每個類別獨立激活每個維度的輸出表示屬于該類別的概率彼此之間不競爭。損失函數也從CrossEntropyLoss換成torch.nn.BCEWithLogitsLoss。這個轉換說難不難但思維方式需要轉一個彎多分類是互相排斥的K選1多標簽是相互獨立的K個二分類。理解了這兩個問題的區別你的分類知識體系就會更完整以后再遇到各種業務場景也能快速判斷用哪種方案。Day18這一天我自己最大的收獲倒不是記住了多少公式而是真正建立了多分類是一個完整閉環的意識從模型輸出層的設計到損失函數的選擇再到評估指標的解讀任何一個環節拿捏不準都會讓整個任務的質量打折扣。如果你現在也卡在二分類往多分類過渡的階段建議你動手寫一個完整的小項目把今天文章里提到的每個環節都親自跑一遍。紙上得來終覺淺等你真的看到自己在CIFAR-10上訓練出了第一個多分類模型那種通了的感覺會比讀十篇筆記都管用。