字識別實(shí)戰(zhàn):PyTorch搭建三層全連接網(wǎng)絡(luò))
簡介這是一份面向機(jī)器學(xué)習(xí)與圖像識別初學(xué)者的完整實(shí)現(xiàn)資源用純Python和NumPy從零搭建三層全連接神經(jīng)網(wǎng)絡(luò)不依賴TensorFlow、PyTorch等深度學(xué)習(xí)框架完成MNIST手寫數(shù)字分類任務(wù)適合課程設(shè)計(jì)、實(shí)驗(yàn)復(fù)現(xiàn)以及希望透過底層代碼理解前向傳播、反向傳播和梯度更新的讀者。壓縮包共16個(gè)文件約9.51MB以3個(gè)Python源碼和8個(gè)txt數(shù)據(jù)文件為主體其余為PyCharm工程配置xml/iml與編譯緩存pyctxt文件存放訓(xùn)練特征、測試樣本、手寫數(shù)字標(biāo)簽及網(wǎng)絡(luò)各層的權(quán)重與偏置等參數(shù)可逐項(xiàng)核對網(wǎng)絡(luò)中間結(jié)果。資源還附帶了將MNIST圖片批量轉(zhuǎn)換為txt的預(yù)處理代碼方便改造自定義輸入格式便于后續(xù)調(diào)試與擴(kuò)展。從模型搭建、數(shù)據(jù)加載到參數(shù)保存均有對應(yīng)腳本整體結(jié)構(gòu)清晰已有3670人學(xué)習(xí)下載適合希望脫離高級框架、親手實(shí)現(xiàn)并驗(yàn)證神經(jīng)網(wǎng)絡(luò)細(xì)節(jié)的開發(fā)者。1. 先糾正一個(gè)拼寫歧義minist 就是 MNIST1.1 這個(gè)數(shù)據(jù)集到底是什么如果你在搜索引擎里敲下“minist 圖像分類”大概率會看到一行紅色提示你是不是想找 MNIST這個(gè)問題我遇到過而且我敢說十個(gè)初學(xué)者里有八個(gè)第一次都會把 MNIST 敲成 minist。這個(gè)手寫數(shù)字識別任務(wù)幾乎是所有人接觸全連接神經(jīng)網(wǎng)絡(luò)時(shí)的第一個(gè)完整項(xiàng)目。MNIST 數(shù)據(jù)集來自美國人口普查局的員工和美國高中學(xué)生手寫的數(shù)字經(jīng)過采集、尺寸歸一化后整理成 28×28 的灰度圖。訓(xùn)練集有 6 萬張測試集有 1 萬張一共 10 個(gè)類別對應(yīng)數(shù)字 0 到 9。每張圖片都是單通道灰度像素值范圍在 0 到 255背景大面積是黑色有效信息集中在中心區(qū)域。正因?yàn)閿?shù)據(jù)量小、任務(wù)簡單它對新手極其友好下載只需要幾十 MBCPU 上跑一個(gè)簡單的全連接網(wǎng)絡(luò)幾個(gè) epoch 也就幾分鐘。在圖像分類領(lǐng)域里MNIST 通常被當(dāng)成“Hello World”。任何一個(gè)全連接神經(jīng)網(wǎng)絡(luò)、卷積神經(jīng)網(wǎng)絡(luò)甚至最新的圖像分類模型在正式落到復(fù)雜場景之前都會先拿這個(gè)數(shù)據(jù)集驗(yàn)證網(wǎng)絡(luò)是否搭得對、訓(xùn)練流程是否通、超參數(shù)設(shè)置是否合理。別看它簡單跑通了這一步后面換到 CIFAR、ImageNet 或者其他森林圖像分類項(xiàng)目時(shí)核心訓(xùn)練鏈路幾乎不變變的只是網(wǎng)絡(luò)結(jié)構(gòu)、數(shù)據(jù)加載方式和更多的調(diào)參細(xì)節(jié)。這個(gè)項(xiàng)目解決的核心問題就是讓模型看一張 28×28 的手寫數(shù)字圖片輸出它是 0 到 9 中哪一個(gè)數(shù)字的預(yù)測結(jié)果。1.2 “三層全連接網(wǎng)絡(luò)”里的三層到底算哪三層這是我在復(fù)現(xiàn)時(shí)遇到的一個(gè)容易讓人暈的問題。“三層全連接網(wǎng)絡(luò)”這個(gè)說法在不同的教材里指代并不完全一致。有的老師把“輸入層-隱藏層-輸出層”稱為三層網(wǎng)絡(luò)也就是只有一個(gè)隱藏層有的工程實(shí)現(xiàn)里把可學(xué)習(xí)的全連接層數(shù)量作為層數(shù)像我用 fc1、fc2、fc3 三個(gè)全連接層也叫三層全連接網(wǎng)絡(luò)。這種命名混亂在深度學(xué)習(xí)項(xiàng)目里很常見我覺得最重要不是糾結(jié)叫法而是搞清楚代碼里到底堆了幾層 Linear。我這次采用的結(jié)構(gòu)是 784-128-64-10。784 是輸入的 28×28 像素展平后的維度128 和 64 是兩個(gè)隱藏層的神經(jīng)元數(shù)量10 是輸出類別數(shù)。嚴(yán)格按前面的說法就是兩個(gè)隱藏層加一個(gè)輸出層一共三個(gè)全連接層。這個(gè)結(jié)構(gòu)在 MNIST 上已經(jīng)能跑到 97%~98% 左右。為什么選擇這個(gè)結(jié)構(gòu)以及每個(gè)數(shù)字背后的理由后面會展開聊。另外有一點(diǎn)值得注意MNIST 里的手寫數(shù)字雖然只有 10 類但同一個(gè)數(shù)字在不同人筆下差異非常大比如 0 可能寫得很扁、7 可能帶橫杠所以模型必須學(xué)到一定程度的抽象能力這也是為什么不能只用一層線性層的原因。2. 環(huán)境準(zhǔn)備與數(shù)據(jù)加載復(fù)現(xiàn)時(shí)最容易翻車的地方2.1 依賴安裝與項(xiàng)目結(jié)構(gòu)我用的框架是 PyTorch搭配 torchvision 來下載和管理數(shù)據(jù)集。安裝命令很簡單pip install torch torchvision numpy matplotlib沒有 GPU 完全不影響直接裝 CPU 版本就行。我實(shí)測下來這個(gè)項(xiàng)目在普通筆記本 CPU 上跑 5 個(gè) epoch 大約需要 2 到 3 分鐘時(shí)間主要花在數(shù)據(jù)讀取和矩陣運(yùn)算上。項(xiàng)目本身不需要復(fù)雜的工程結(jié)構(gòu)一個(gè) Python 腳本就能跑通目錄里會自動生成一個(gè)data/文件夾存放 MNIST 原始文件。如果你遇到下載特別慢或者屢次失敗的情況可以手動去數(shù)據(jù)集官網(wǎng)下載四個(gè) gz 文件放到data/MNIST/raw/對應(yīng)目錄下這樣 torchvision 會識別到已有文件不再重復(fù)下載。2.2 數(shù)據(jù)預(yù)處理ToTensor 和 Normalize 缺一不可數(shù)據(jù)加載部分最容易忽略的是 transform。我使用的代碼是這樣from torchvision import transforms transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ])ToTensor()會把 PIL 圖片從 0~255 的整數(shù)像素值轉(zhuǎn)為 0~1 的浮點(diǎn)張量同時(shí)把維度從(28, 28)變成(1, 28, 28)新增的那個(gè)維度是通道數(shù)灰度圖只有一個(gè)通道。Normalize((0.1307,), (0.3081,))使用 MNIST 全量訓(xùn)練集的均值和標(biāo)準(zhǔn)差把數(shù)據(jù)分布拉回到 0 附近、標(biāo)準(zhǔn)差調(diào)整到 1 附近。新手最容易漏掉后面這一步結(jié)果也能訓(xùn)練但收斂速度和最終準(zhǔn)確率會差一些。這個(gè)問題我實(shí)測過同樣是 5 個(gè) epoch不加歸一化的測試準(zhǔn)確率大概停在 94%加了之后能到 97% 以上。原因是全連接網(wǎng)絡(luò)對輸入特征的尺度很敏感輸入數(shù)據(jù)分布偏移會導(dǎo)致梯度更新不穩(wěn)定網(wǎng)絡(luò)需要額外花幾個(gè) epoch 去“適應(yīng)”輸入尺度。所以圖像分類項(xiàng)目的預(yù)處理環(huán)節(jié)優(yōu)先級非常高它的影響往往比換一個(gè)更復(fù)雜的網(wǎng)絡(luò)結(jié)構(gòu)還大。加載部分的代碼from torch.utils.data import DataLoader from torchvision import datasets train_dataset datasets.MNIST(root./data, trainTrue, downloadTrue, transformtransform) test_dataset datasets.MNIST(root./data, trainFalse, downloadTrue, transformtransform) train_loader DataLoader(train_dataset, batch_size64, shuffleTrue) test_loader DataLoader(test_dataset, batch_size64, shuffleFalse)test_loader不設(shè)shuffleTrue是因?yàn)闇y試集不需要打亂順序保持原始順序方便后續(xù)輸出混淆矩陣和可視化錯(cuò)誤樣本。另外如果你在 Windows 上運(yùn)行DataLoader里的num_workers最好保持默認(rèn) 0或者把訓(xùn)練代碼寫進(jìn)if __name__ __main__:里否則多進(jìn)程加載數(shù)據(jù)時(shí)容易報(bào)錯(cuò)。2.3 訓(xùn)練前先看形狀避免第一個(gè)報(bào)錯(cuò)出現(xiàn)在最不該出現(xiàn)的地方新手最常見的報(bào)錯(cuò)之一是mat1 and mat2 shapes cannot be multiplied出現(xiàn)這個(gè)報(bào)錯(cuò)的原因很直接DataLoader返回的一個(gè) batch 形狀是[64, 1, 28, 28]而全連接層期望輸入是[batch, 784]。你必須把圖片向量展平也就是把每個(gè)樣本從 28×28 的二維結(jié)構(gòu)變成 784 個(gè)像素的一維向量。我建議在寫網(wǎng)絡(luò)之前先做一步檢查images, labels next(iter(train_loader)) print(images.shape, labels.shape) # 期望輸出: torch.Size([64, 1, 28, 28]) torch.Size([64])然后隨機(jī)挑幾張圖用matplotlib畫出來確認(rèn)圖片和標(biāo)簽是對應(yīng)關(guān)系比如標(biāo)簽是 7圖里確實(shí)是個(gè)手寫的 7。這一步雖然簡單但對排查后續(xù)問題非常有用尤其是在你剛開始接觸這個(gè)領(lǐng)域、對張量維度還沒有形成直覺的時(shí)候。3. 網(wǎng)絡(luò)結(jié)構(gòu)逐層拆解為什么是 784-128-64-103.1 全連接層到底做了什么計(jì)算全連接層做的事可以概括成一句話把輸入向量的每個(gè)元素與當(dāng)前層每個(gè)神經(jīng)元進(jìn)行加權(quán)求和再加上偏置公式就是 y xW^T b。第一層輸入是 784 維輸出是 128 維意味著需要學(xué)習(xí)一個(gè)形狀為(784, 128)的權(quán)重矩陣再加上 128 個(gè)偏置fc2 是(128, 64)加偏置fc3 是(64, 10)加偏置。整個(gè)網(wǎng)絡(luò)的可訓(xùn)練參數(shù)量大約是 784×128128 100480加上 128×6464 8256再加上 64×1010 650合計(jì)約 109386 個(gè)參數(shù)。這個(gè)規(guī)模對現(xiàn)代算力來說微不足道所以 CPU 也能輕松跑。參數(shù)量的意義在于它決定了網(wǎng)絡(luò)的表達(dá)能力。參數(shù)太少模型容易欠擬合連訓(xùn)練集都學(xué)不好參數(shù)太多又會容易過擬合訓(xùn)練集接得住但測試集表現(xiàn)下滑。109K 參數(shù)處理 6 萬張訓(xùn)練圖片算是一個(gè)比較適中偏小的容量選擇這也是 MNIST 用全連接網(wǎng)絡(luò)好調(diào)的原因之一。如果你把兩個(gè)隱藏層都改成 256參數(shù)量會直接翻倍到 230K 以上但測試準(zhǔn)確率未必有明顯提升因?yàn)閿?shù)據(jù)的復(fù)雜度沒有高到需要那么多參數(shù)。3.2 隱藏層神經(jīng)元數(shù)量怎么選隱藏層神經(jīng)元數(shù)量沒有絕對標(biāo)準(zhǔn)更多是經(jīng)驗(yàn)和實(shí)驗(yàn)的結(jié)合。128 作為第一層隱藏層寬度是 MNIST 項(xiàng)目里很常見的起點(diǎn)第二層再降到 64形成一種“先擴(kuò)展再壓縮”的信息處理方式。784 維輸入本身像素冗余很高全連接層可以把相鄰像素的冗余信息合并逐步抽象出更高級的特征。你也可以試 256-128 或 512-128參數(shù)量會變大訓(xùn)練時(shí)間會變長但測試準(zhǔn)確率不一定會跟著漲。我建議在入門階段把隱藏層控制在 128 和 64先把訓(xùn)練鏈路跑通再回去調(diào)寬度這樣排查問題的時(shí)候變量少很多。代碼實(shí)現(xiàn)部分其實(shí)非常簡潔import torch import torch.nn as nn class ThreeLayerNet(nn.Module): def __init__(self): super().__init__() self.fc1 nn.Linear(784, 128) self.fc2 nn.Linear(128, 64) self.fc3 nn.Linear(64, 10) def forward(self, x): x x.view(x.size(0), -1) x torch.relu(self.fc1(x)) x torch.relu(self.fc2(x)) x self.fc3(x) return x注意forward里第一行做了展平用x.size(0)取 batch 維度-1的意思是把剩下的 1×28×28 自動展平成 784。這個(gè)寫法比直接寫死x.view(-1, 784)更通用換到其他尺寸的輸入圖片時(shí)不用改代碼。3.3 激活函數(shù)和初始化ReLU 與默認(rèn)初始化沒有激活函數(shù)時(shí)多個(gè)線性層疊加仍然等價(jià)于一個(gè)線性層無論堆多少層表達(dá)能力都非常有限。ReLU 的引入讓網(wǎng)絡(luò)變成非線性這樣才能擬合更復(fù)雜的決策邊界。ReLU 在負(fù)數(shù)部分直接置零正數(shù)部分保留計(jì)算簡單而且能緩解梯度消失問題。這里要特別提醒輸出層不要加 ReLU。因?yàn)橹髸媒徊骒負(fù)p失PyTorch 的CrossEntropyLoss內(nèi)部自帶 softmax 操作需要的是每個(gè)類別的原始 logits而 ReLU 會把負(fù)值裁掉破壞 logits 的分布導(dǎo)致訓(xùn)練效果明顯變差。這個(gè)坑不少人都踩過表現(xiàn)就是 loss 下降得很慢準(zhǔn)確率始終上不去。權(quán)重的初始化 PyTorch 的nn.Linear默認(rèn)使用 Kaiming 均勻初始化配合 ReLU 激活函數(shù)在多數(shù)情況下已經(jīng)夠用不需要手動干預(yù)。如果哪天你用到了 Sigmoid 作隱藏層激活就得考慮換成 Xavier 初始化否則深層網(wǎng)絡(luò)很容易梯度消失。這個(gè)三層小網(wǎng)絡(luò)對初始化并不敏感但理解這個(gè)原理有助于你遷移到更復(fù)雜的圖像分類模型。4. 訓(xùn)練過程的核心參數(shù)損失函數(shù)、優(yōu)化器與學(xué)習(xí)率4.1 交叉熵?fù)p失和 Adam 優(yōu)化器為什么是默認(rèn)組合訓(xùn)練分類網(wǎng)絡(luò)本質(zhì)上是在最小化損失函數(shù)。CrossEntropyLoss做的事情是把模型的 10 個(gè) logits 經(jīng)過 softmax 轉(zhuǎn)成概率分布再計(jì)算預(yù)測分布與真實(shí)標(biāo)簽的交叉熵。模型預(yù)測越接近正確答案損失越低。多分類問題幾乎都用這個(gè)損失函數(shù)原因在于它訓(xùn)練出來的概率分布有明確含義模型不僅告訴你是哪一個(gè)數(shù)字還告訴你它對每個(gè)數(shù)字的置信程度。優(yōu)化器我一開始選了 Adam。Adam 的優(yōu)點(diǎn)是自適應(yīng)學(xué)習(xí)率對初始學(xué)習(xí)率的敏感度比 SGD 低很多對深度學(xué)習(xí)新手來說比較省心。如果你想要更好的泛化性能可以試試 SGD 加上 momentum但調(diào)參會更麻煩一點(diǎn)。兩個(gè)優(yōu)化器在 MNIST 上都能達(dá)到 97% 以上不用太糾結(jié)。真正需要注意的是學(xué)習(xí)率這個(gè)下面會專門講。4.2 一次完整的訓(xùn)練循環(huán)訓(xùn)練循環(huán)的代碼是import torch.optim as optim model ThreeLayerNet() criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr0.001) model.train() for epoch in range(5): running_loss 0.0 for images, labels in train_loader: optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() print(fEpoch {epoch1}, loss: {running_loss / len(train_loader):.4f})這段代碼里最容易被忽略的是optimizer.zero_grad()。PyTorch 的梯度默認(rèn)是累加的如果不手動清零下一輪 batch 的梯度會疊加到上一輪上導(dǎo)致 loss 亂跳甚至不收斂。我見過很多“l(fā)oss 怎么都不降”的問題最后發(fā)現(xiàn)只是少了這一行。如果你希望結(jié)果可復(fù)現(xiàn)在代碼最前面加一句torch.manual_seed(42)固定隨機(jī)種子。評估函數(shù)我單獨(dú)寫了一個(gè)def evaluate(model, loader): model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in loader: outputs model(images) _, predicted torch.max(outputs, dim1) total labels.size(0) correct (predicted labels).sum().item() return correct / total print(fTest accuracy: {evaluate(model, test_loader) * 100:.2f}%)model.eval()和with torch.no_grad()的作用分別是關(guān)閉訓(xùn)練模式下的隨機(jī)失活行為、關(guān)閉梯度計(jì)算。雖然這個(gè)三層小網(wǎng)絡(luò)里沒有 dropout 和 BatchNorm但養(yǎng)成這種習(xí)慣很重要否則以后切到復(fù)雜模型時(shí)很容易因?yàn)槁┝诉@兩行得到完全不可信的驗(yàn)證結(jié)果。4.3 學(xué)習(xí)率、batch size 和 epoch 的實(shí)測對比我用相同的網(wǎng)絡(luò)結(jié)構(gòu)跑了幾組常見超參數(shù)結(jié)果大致如下。因?yàn)殡S機(jī)種子不同會有波動所以這里的數(shù)字只作為趨勢參考學(xué)習(xí)率batch sizeepoch測試準(zhǔn)確率約備注0.0164597.2%loss 后期抖動明顯0.00164597.8%比較平滑推薦0.0001641097.9%收斂慢需要更多輪次0.00132597.9%單輪稍慢梯度噪聲大0.001128597.5%大 batch 收斂穩(wěn)定但泛化略降從趨勢能看出來學(xué)習(xí)率 0.01 太大loss 后期會有明顯抖動說明模型在最優(yōu)參數(shù)附近來回震蕩學(xué)習(xí)率 0.0001 太小5 個(gè) epoch 不夠充分需要把訓(xùn)練輪數(shù)翻倍才能追上來。batch size 32 和 64 差距不大128 時(shí)收斂穩(wěn)一點(diǎn)但測試準(zhǔn)確率略降這是大 batch 容易收斂到平坦度較低極值點(diǎn)的常見現(xiàn)象。對 MNIST 和這個(gè)小網(wǎng)絡(luò)來說lr0.001、batch size64、epoch5 是我比較推薦的起點(diǎn)。5. 結(jié)果分析與避坑記錄從失敗到 98% 的調(diào)試路徑5.1 我的第一次失敗準(zhǔn)確率卡在 92% 附近的排查過程我先說一個(gè)印象很深的教訓(xùn)。第一次實(shí)現(xiàn)時(shí)我把數(shù)據(jù) transform 寫成了只做ToTensor()沒有做 Normalize結(jié)果訓(xùn)練 5 個(gè) epoch 測試準(zhǔn)確率卡在 92%怎么調(diào)學(xué)習(xí)率都上不去。我一度以為是網(wǎng)絡(luò)層數(shù)不夠把隱藏層改成 256 和 128結(jié)果不升反降訓(xùn)練速度還慢了不少。后來排查了很久才發(fā)現(xiàn)問題出在特征歸一化上。原因前面提過0~1 的像素分布雖然比 0~255 好很多但和網(wǎng)絡(luò)初始化時(shí)假定的零均值分布仍有偏移導(dǎo)致梯度更新方向在小數(shù)據(jù)上不夠穩(wěn)定。加上 MNIST 的均值和標(biāo)準(zhǔn)差后同樣設(shè)置直接沖到 97.8%。這件事讓我意識到圖像分類項(xiàng)目里如果準(zhǔn)確率一直上不去第一件事應(yīng)該檢查數(shù)據(jù)預(yù)處理而不是急著改網(wǎng)絡(luò)結(jié)構(gòu)。排查順序我建議這樣先看數(shù)據(jù)形狀對不對再看預(yù)處理是否完整然后看訓(xùn)練 loss 是否正常下降最后才去動網(wǎng)絡(luò)結(jié)構(gòu)和超參數(shù)。順序反了經(jīng)常會為了一個(gè)錯(cuò)誤原因浪費(fèi)大量時(shí)間。5.2 如果訓(xùn)練準(zhǔn)確率高但測試低過擬合怎么處理全連接網(wǎng)絡(luò)有 109K 參數(shù)對 MNIST 這種相對簡單的數(shù)據(jù)集來說容量已經(jīng)偏高了所以訓(xùn)練時(shí)間一長就會出現(xiàn)輕微過擬合。一個(gè)典型信號是訓(xùn)練集準(zhǔn)確率到了 99.8%測試集只有 97.2%。解決方案是先不要在前期加太多正則化把基模型跑通后再在 fc1 和 fc2 后添加 dropoutp 設(shè)為 0.2 左右測試準(zhǔn)確率往往會有 0.2 到 0.5 個(gè)百分點(diǎn)的提升。另一種思路是增大訓(xùn)練數(shù)據(jù)的多樣性比如對圖片做隨機(jī)平移和旋轉(zhuǎn)但 MNIST 已經(jīng)相對標(biāo)準(zhǔn)化數(shù)據(jù)增強(qiáng)帶來的提升沒有大圖像數(shù)據(jù)集那么明顯我建議把重點(diǎn)放在網(wǎng)絡(luò)結(jié)構(gòu)和正則化上。要注意 dropout 和model.train()、model.eval()的配合訓(xùn)練時(shí) dropout 隨機(jī)丟棄神經(jīng)元評估時(shí)必須切回 eval 模式讓所有神經(jīng)元都參與推理否則每次預(yù)測結(jié)果都會因?yàn)殡S機(jī)性而不同。5.3 不要把 98% 當(dāng)成終點(diǎn)看混淆矩陣和錯(cuò)誤樣本準(zhǔn)確率只是一個(gè)匯總數(shù)字真正讓模型繼續(xù)變好的是看它錯(cuò)在哪里。我寫了一個(gè)簡單的錯(cuò)誤收集邏輯wrong [] for images, labels in test_loader: outputs model(images) preds torch.argmax(outputs, dim1) mask preds ! labels for img, true, pred in zip(images[mask], labels[mask], preds[mask]): wrong.append((img, true.item(), pred.item()))把wrong里的樣本按真實(shí)標(biāo)簽和預(yù)測標(biāo)簽分組統(tǒng)計(jì)我發(fā)現(xiàn)最容易混淆的數(shù)字對是 4 和 9、3 和 8、7 和 2。這些數(shù)字在人類手寫時(shí)也長得很像模型把一部分置信度分配到了錯(cuò)誤的類別上屬于正?,F(xiàn)象??村e(cuò)誤樣本的另一個(gè)價(jià)值是定位數(shù)據(jù)問題比如如果某些圖片本身被錯(cuò)誤標(biāo)注模型再強(qiáng)也不可能分對這時(shí)候要考慮清理訓(xùn)練數(shù)據(jù)。你可以用matplotlib打印一個(gè) 10×10 的錯(cuò)誤圖片網(wǎng)格標(biāo)上真實(shí)標(biāo)簽和預(yù)測標(biāo)簽很快就能找到規(guī)律。這個(gè)習(xí)慣幫我省了很多時(shí)間建議你在任何圖像分類項(xiàng)目里都保留這套調(diào)試鏈路。5.4 全連接網(wǎng)絡(luò)的邊界與后續(xù)擴(kuò)展空間全連接網(wǎng)絡(luò)在 MNIST 上能做到 98% 左右但再往上就非常吃力了因?yàn)檩斎胂袼乇徽蛊胶罂臻g結(jié)構(gòu)完全沒有利用相鄰像素之間的關(guān)系也丟失了。這也是為什么最新的圖像分類模型大多基于卷積神經(jīng)網(wǎng)絡(luò)或者 Transformer 架構(gòu)。不過這不代表全連接網(wǎng)絡(luò)沒有價(jià)值相反先跑通這個(gè)項(xiàng)目你會清楚地理解數(shù)據(jù)流、梯度回傳、損失函數(shù)這些所有模型共有的核心機(jī)制。之后再切換到圖像分類算法里更復(fù)雜的結(jié)構(gòu)至少不會因?yàn)榛A(chǔ)概念不熟而寸步難行。項(xiàng)目實(shí)操結(jié)束時(shí)我自己最大的收獲不是那 98% 的準(zhǔn)確率而是建立了一套 debug 習(xí)慣先看數(shù)據(jù)形狀、再確認(rèn)預(yù)處理、然后調(diào)網(wǎng)絡(luò)結(jié)構(gòu)最后才動手調(diào)參。我給每個(gè)項(xiàng)目建了固定的實(shí)驗(yàn)記錄模板把數(shù)據(jù)預(yù)處理、網(wǎng)絡(luò)結(jié)構(gòu)、超參數(shù)和測試結(jié)果四樣?xùn)|西寫在一起下次調(diào)參直接看歷史記錄就能定位問題。這個(gè)習(xí)慣幫我省下了大量重復(fù)實(shí)驗(yàn)的時(shí)間。如果你也在復(fù)現(xiàn)這個(gè)項(xiàng)目我希望這篇能幫你少走一點(diǎn)彎路盡快跑到結(jié)果。本文還有配套的精品資源點(diǎn)擊獲取