數(shù)字識(shí)別實(shí)戰(zhàn):從MNIST數(shù)據(jù)到99%準(zhǔn)確率)
簡(jiǎn)介卷積神經(jīng)網(wǎng)絡(luò)CNN是深度學(xué)習(xí)中處理圖像分類任務(wù)的核心技術(shù)之一它通過(guò)卷積、池化與全連接層的協(xié)同工作自動(dòng)從原始像素中提取局部特征實(shí)現(xiàn)對(duì)圖像的高效識(shí)別。與普通全連接網(wǎng)絡(luò)相比CNN具備參數(shù)共享和平移不變性在圖像任務(wù)中泛化能力更強(qiáng)因此成為計(jì)算機(jī)視覺(jué)領(lǐng)域的基礎(chǔ)模型。從技術(shù)價(jià)值來(lái)看CNN不僅能用于經(jīng)典的手寫(xiě)數(shù)字識(shí)別還能遷移到物體檢測(cè)、人臉識(shí)別等復(fù)雜場(chǎng)景是工程實(shí)踐中高頻使用的模型架構(gòu)。在模型訓(xùn)練過(guò)程中選擇合適的深度學(xué)習(xí)框架至關(guān)重要PyTorch憑借動(dòng)態(tài)計(jì)算圖和靈活的調(diào)試體驗(yàn)深受開(kāi)發(fā)者喜愛(ài)。本文以MNIST數(shù)據(jù)集為例完整講解從數(shù)據(jù)加載、數(shù)據(jù)歸一化、DataLoader批處理、CNN模型搭建到訓(xùn)練評(píng)估與可視化的全流程并最終達(dá)到99%以上的測(cè)試準(zhǔn)確率為課程設(shè)計(jì)、畢業(yè)設(shè)計(jì)以及入門(mén)深度學(xué)習(xí)工程實(shí)踐提供可復(fù)現(xiàn)的參考路徑。 作為一名過(guò)來(lái)人我太清楚畢業(yè)設(shè)計(jì)或期末大作業(yè)最怕的不是不會(huì)寫(xiě)代碼而是拿到一個(gè)題目后不知道從哪下手。手寫(xiě)數(shù)字識(shí)別這個(gè)方向是很多同學(xué)的首選因?yàn)樗扔凶銐虻摹凹夹g(shù)含量”又不會(huì)難到無(wú)法收尾。這篇內(nèi)容我打算把整個(gè)項(xiàng)目從零到一拆開(kāi)揉碎講清楚——基于 Python 實(shí)現(xiàn) CNN 卷積神經(jīng)網(wǎng)絡(luò)完成手寫(xiě)數(shù)字識(shí)別配套完整源碼、詳細(xì)注釋和數(shù)據(jù)集處理方案。不管你是準(zhǔn)備交期末作業(yè)還是畢業(yè)論文需要實(shí)驗(yàn)章節(jié)這份實(shí)操路線都可以直接參考。我默認(rèn)你已經(jīng)具備一點(diǎn)點(diǎn) Python 語(yǔ)法基礎(chǔ)但不需要會(huì)復(fù)雜的數(shù)學(xué)推導(dǎo)。CNN 里那些卷積、池化、全連接的概念我會(huì)用大白話加代碼一起講。項(xiàng)目里我會(huì)用 PyTorch 作為深度學(xué)習(xí)框架因?yàn)樗{(diào)試直觀、寫(xiě)起來(lái)靈活而且學(xué)術(shù)界和工業(yè)界都在用答辯時(shí)老師不會(huì)挑框架的毛病。最終會(huì)在 MNIST 數(shù)據(jù)集上跑出 99% 以上的測(cè)試準(zhǔn)確率這個(gè)指標(biāo)對(duì)于課程設(shè)計(jì)和本科畢設(shè)已經(jīng)完全夠用了。1. 項(xiàng)目定調(diào)畢設(shè)級(jí) CNN 手寫(xiě)數(shù)字識(shí)別怎么規(guī)劃1.1 為什么手寫(xiě)數(shù)字識(shí)別適合作為課程設(shè)計(jì)和畢業(yè)設(shè)計(jì)選題手寫(xiě)數(shù)字識(shí)別本質(zhì)上是圖像分類任務(wù)輸入是一張 28×28 的灰度圖片輸出是 0 到 9 這十個(gè)數(shù)字中的某一個(gè)類別。這個(gè)任務(wù)看起來(lái)簡(jiǎn)單但它把深度學(xué)習(xí)最核心的流程全部涵蓋了數(shù)據(jù)加載、模型搭建、訓(xùn)練調(diào)參、評(píng)估分析、結(jié)果可視化。評(píng)閱老師拿到一份項(xiàng)目看到你把這五個(gè)環(huán)節(jié)都完整實(shí)現(xiàn)了印象分天然就會(huì)高。選 MNIST 數(shù)據(jù)集還有幾個(gè)現(xiàn)實(shí)原因。第一是數(shù)據(jù)規(guī)模適中6 萬(wàn)張訓(xùn)練圖片加 1 萬(wàn)張測(cè)試圖片在我的筆記本 CPU 上跑完 10 個(gè) epoch 也就幾分鐘完全不需要依賴昂貴的 GPU。第二是圖片分辨率低28×28 單通道意味著模型結(jié)構(gòu)可以很輕量即使把網(wǎng)絡(luò)層數(shù)加深一點(diǎn)參數(shù)量仍然在可控范圍內(nèi)。第三是生態(tài)成熟數(shù)據(jù)下載、預(yù)處理、效果對(duì)比都有一套標(biāo)準(zhǔn)參考不太會(huì)出現(xiàn)你復(fù)現(xiàn)不出別人結(jié)果的情況。很多同學(xué)會(huì)糾結(jié)“這么經(jīng)典的任務(wù)會(huì)不會(huì)太簡(jiǎn)單體現(xiàn)不出水平”。我的看法是你能不能在有限時(shí)間內(nèi)把經(jīng)典任務(wù)做完整、講清楚比用花哨模型堆砌重要的多。導(dǎo)師真正在乎的是你有沒(méi)有理解卷積神經(jīng)網(wǎng)絡(luò)在做什么而不是你是不是用了一個(gè)冷門(mén)數(shù)據(jù)集。后續(xù)你想加分完全可以在改進(jìn)部分加入數(shù)據(jù)增強(qiáng)、模型結(jié)構(gòu)調(diào)優(yōu)甚至部署成 Web 應(yīng)用這些都是可擴(kuò)展的點(diǎn)。1.2 深度學(xué)習(xí)框架選型PyTorch 還是 TensorFlow寫(xiě)這段的時(shí)候我其實(shí)很有感觸。我最早學(xué)的是 TensorFlow 1.x那時(shí)候 API 設(shè)計(jì)反人類每次搭模型都要先畫(huà)計(jì)算圖調(diào)試一個(gè)維度錯(cuò)誤能卡一個(gè)下午。后來(lái) PyTorch 起來(lái)了它的動(dòng)態(tài)計(jì)算圖機(jī)制對(duì)新手極其友好——你寫(xiě)代碼的方式和程序?qū)嶋H執(zhí)行的方式一致報(bào)錯(cuò)也能直接定位到 Python 代碼行不需要繞一層抽象。所以我強(qiáng)烈建議畢設(shè)項(xiàng)目用 PyTorch。PyTorch 的生態(tài)也很完善torchvision 提供了 MNIST 數(shù)據(jù)集的直接下載接口torch.nn 里面卷積層、池化層、全連接層都是現(xiàn)成的模塊。你用 $cnn$ 結(jié)構(gòu)搭建一個(gè)模型核心代碼不會(huì)超過(guò) 50 行。這在答辯場(chǎng)景下是優(yōu)勢(shì)因?yàn)槔蠋熌弥愕脑创a逐行問(wèn)你能快速說(shuō)清楚每一行的作用而不是搬出一大堆框架自動(dòng)生成的東西。當(dāng)然我也承認(rèn)如果團(tuán)隊(duì)或者學(xué)校課程一直用 TensorFlow/Keras那就沒(méi)必要強(qiáng)行換。衡量標(biāo)準(zhǔn)只有一個(gè)你能否在截止日期前獨(dú)立完成閉環(huán)。如果你對(duì) PyTorch 完全零基礎(chǔ)但會(huì)用 Python大概需要兩到三天時(shí)間適應(yīng)它的數(shù)據(jù)流和訓(xùn)練循環(huán)寫(xiě)法。這個(gè)時(shí)間成本放在期末周里不算小所以選型要趁早。1.3 代碼結(jié)構(gòu)與文件規(guī)劃項(xiàng)目不要把所有代碼堆在一個(gè) notebook 里雖然 Jupyter Notebook 適合演示但作為交付源碼還是建議拆分成模塊化文件。我最終的目錄結(jié)構(gòu)大致如下mnist_cnn/ ├── data/ │ └── MNIST/ # 數(shù)據(jù)集存放位置自動(dòng)下載 ├── models/ │ └── model.py # CNN 網(wǎng)絡(luò)結(jié)構(gòu)定義 ├── utils/ │ ├── dataset.py # 數(shù)據(jù)加載與預(yù)處理 │ ├── train.py # 訓(xùn)練邏輯 │ └── visualize.py # 訓(xùn)練曲線、混淆矩陣可視化 ├── main.py # 一鍵運(yùn)行數(shù)據(jù) → 訓(xùn)練 → 評(píng)估 ├── predict.py # 單張圖片推理演示 ├── requirements.txt └── README.md我堅(jiān)持模塊化的原因有兩點(diǎn)。第一是可維護(hù)性好你想調(diào)整網(wǎng)絡(luò)結(jié)構(gòu)只改 model.py想換數(shù)據(jù)增強(qiáng)策略只改 dataset.py互不影響。第二是答辯答辯時(shí)你可以講清楚軟件工程思想這也是老師常問(wèn)的問(wèn)題點(diǎn)。很多同學(xué)在準(zhǔn)備期末大作業(yè)時(shí)能力完全夠但代碼亂成一鍋粥最后扣分很冤。模塊化即使不加分也絕不會(huì)扣分。2. 數(shù)據(jù)處理與加載從 MNIST 下載到 DataLoader2.1 MNIST 數(shù)據(jù)集的基本情況MNIST 全稱是 Modified National Institute of Standards and Technology手寫(xiě)數(shù)字?jǐn)?shù)據(jù)集的經(jīng)典中的經(jīng)典。它包含 0 到 9 十個(gè)類別每張圖片為 28 像素寬、28 像素高、單通道灰度圖像素值范圍在 0 到 255 之間0 表示黑色背景255 表示白色筆跡。訓(xùn)練集有 60000 張測(cè)試集有 10000 張。注意訓(xùn)練集和測(cè)試集是官方劃分好的我們?cè)谧鰧?shí)驗(yàn)時(shí)千萬(wàn)不能把自己的驗(yàn)證集從測(cè)試集里切否則測(cè)試集就失去了“沒(méi)有見(jiàn)過(guò)的數(shù)據(jù)”的意義。在 PyTorch 中用 torchvision.datasets.MNIST 接口下載時(shí)可以看到 train 參數(shù)trainTrue 下載訓(xùn)練集trainFalse 下載測(cè)試集。這里有一個(gè)被很多新手忽略的點(diǎn)你直接拿到的圖片是 PIL 格式不是 PyTorch 能直接計(jì)算的張量。所以每次取出一條數(shù)據(jù)都要先完成格式轉(zhuǎn)換最常見(jiàn)的手段就是使用 torchvision.transforms.ToTensor()它會(huì)把 PIL 圖片轉(zhuǎn)成形狀為 (C, H, W) 的張量并把像素值從 0~255 縮放到 0~1。2.2 像素歸一化為什么要除以 255我知道有些同學(xué)會(huì)偷懶不做歸一化直接把 0~255 的像素塞進(jìn)網(wǎng)絡(luò)結(jié)果訓(xùn)練時(shí)發(fā)現(xiàn) loss 很難下降。原因并不神秘。神經(jīng)網(wǎng)絡(luò)中每一層的參數(shù)更新依賴梯度的反向傳播如果輸入特征數(shù)值范圍過(guò)大會(huì)導(dǎo)致某些層的加權(quán)求和結(jié)果很大激活函數(shù)進(jìn)入飽和區(qū)梯度接近于零參數(shù)幾乎無(wú)法更新。把像素值歸一化到 0~1 或者更常見(jiàn)的零均值單位方差后模型收斂速度會(huì)有肉眼可見(jiàn)的提升。torchvision.transforms.ToTensor() 內(nèi)部已經(jīng)替你做了除以 255 的操作所以只要你用了這個(gè) transform輸入到模型的數(shù)據(jù)范圍就是 [0,1]。你還可以再加 Normalize((0.1307,), (0.3081,))這兩個(gè)值是 MNIST 數(shù)據(jù)集的全局均值和標(biāo)準(zhǔn)差在社區(qū)里已經(jīng)是公開(kāi)基礎(chǔ)信息。作用是把數(shù)據(jù)標(biāo)準(zhǔn)化到均值為 0、標(biāo)準(zhǔn)差為 1 的分布進(jìn)一步幫助訓(xùn)練穩(wěn)定。需要強(qiáng)調(diào)的是Normalize 操作中的均值和標(biāo)準(zhǔn)差必須和數(shù)據(jù)集本身匹配。如果你后面更換了 Fashion-MNIST 等數(shù)據(jù)集這兩個(gè)參數(shù)就要重新計(jì)算不能直接抄過(guò)來(lái)用。我見(jiàn)過(guò)有人把 ImageNet 的均值標(biāo)準(zhǔn)差用到 MNIST 上雖然模型最后也能跑但訓(xùn)練曲線明顯不平滑所以你最好不要這么做。2.3 DataLoader 與 batch 概念數(shù)據(jù)準(zhǔn)備中另一個(gè)核心概念是 batch。為什么要用 batch 而不是一次把 60000 張圖片全都喂進(jìn)去如果你嘗試過(guò)全批量梯度下降就明白顯存會(huì)被瞬間撐爆而且訓(xùn)練過(guò)程中損失下降路徑非常僵硬。反過(guò)來(lái)如果每次只喂一張圖權(quán)重更新方向波動(dòng)太大損失函數(shù)像心電圖一樣上下亂跳收斂效率極低。所以 PyTorch 提供了 DataLoader 工具讓我們按 batch 取數(shù)據(jù)。常見(jiàn)選擇是 batch_size64 或 128。以 64 為例每輪迭代從訓(xùn)練集中隨機(jī)抽取 64 張圖片計(jì)算這 64 張的平均梯度然后更新一次參數(shù)。60000 張圖片全部走過(guò)一遍算一個(gè) epoch一個(gè) epoch 有 938 個(gè)這樣的迭代。DataLoader 還有一個(gè) shuffle 參數(shù)訓(xùn)練階段設(shè)為 True。這個(gè)至關(guān)重要因?yàn)槿绻麛?shù)據(jù)原本按標(biāo)簽順序排列不打亂的話每個(gè) batch 內(nèi)可能全是同一個(gè)數(shù)字模型學(xué)到的特征會(huì)產(chǎn)生嚴(yán)重偏移。測(cè)試階段一般設(shè)為 False因?yàn)闇y(cè)試只用前向傳播不需要考慮梯度更新的隨機(jī)性。# utils/dataset.py 核心代碼 from torchvision import datasets, transforms from torch.utils.data import DataLoader def get_dataloader(batch_size64, use_augmentFalse): transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_dataset datasets.MNIST( root./data, trainTrue, downloadTrue, transformtransform ) test_dataset datasets.MNIST( root./data, trainFalse, downloadTrue, transformtransform ) train_loader DataLoader( train_dataset, batch_sizebatch_size, shuffleTrue ) test_loader DataLoader( test_dataset, batch_sizebatch_size, shuffleFalse ) return train_loader, test_loader上面的代碼中downloadTrue 會(huì)在第一次運(yùn)行時(shí)自動(dòng)下載數(shù)據(jù)到 ./data/MNIST 目錄。網(wǎng)絡(luò)通暢的情況下下載很順利如果反復(fù)失敗你可以找一臺(tái)有網(wǎng)環(huán)境的機(jī)器把文件下載后拷貝過(guò)來(lái)也可以手動(dòng)解壓到指定目錄只要目錄結(jié)構(gòu)符合 torchvision 的預(yù)期即可。具體的排障方法我在后面第七部分單獨(dú)整理一個(gè)清單。3. CNN 模型搭建手寫(xiě)數(shù)字識(shí)別背后的圖像原理3.1 卷積層、池化層、全連接層各司其職現(xiàn)在進(jìn)入重頭戲也就是 CNN 卷積神經(jīng)網(wǎng)絡(luò)本身。我在理解這個(gè)模型時(shí)最有用的類比是“圖像濾鏡”。你可以把卷積層想象成一組可以學(xué)習(xí)的濾鏡每個(gè)濾鏡掃描整張圖片提取一種特定的局部特征比如邊緣、拐角、筆畫(huà)的粗細(xì)。一開(kāi)始網(wǎng)絡(luò)不知道哪些特征重要但是通過(guò)訓(xùn)練數(shù)據(jù)反向傳播濾鏡會(huì)自動(dòng)調(diào)整參數(shù)最終保留有用的特征。卷積運(yùn)算有幾個(gè)關(guān)鍵概念需要解釋。第一是局部感受野每次卷積核只觀察輸入圖上一個(gè)小窗口比如 3×3而不是看整張圖。這符合圖像的天然屬性離得很遠(yuǎn)的像素之間關(guān)聯(lián)性弱沒(méi)必要一開(kāi)始就讓它們直接相連。第二是參數(shù)共享同一個(gè)卷積核掃過(guò)整張圖所有位置時(shí)權(quán)重相同。這大大減少了模型參數(shù)量也賦予網(wǎng)絡(luò)平移不變性也就是說(shuō)一個(gè)數(shù)字出現(xiàn)在圖片左上角還是右下角都能被同一個(gè)特征提取器識(shí)別。池化層的作用是降維最常用的是最大池化把 2×2 窗口中的最大值選出來(lái)。這樣做一方面縮小了特征圖的尺寸減少了后續(xù)計(jì)算量另一方面保留了最有響應(yīng)強(qiáng)度的特征讓模型對(duì)輕微位移和形變更加魯棒。我在這個(gè)項(xiàng)目中用了兩個(gè)卷積塊加兩個(gè)池化層圖片從 28×28 逐漸變成 14×14再變成 7×7特征通道從 1 擴(kuò)到 32 再擴(kuò)到 64。通道變多意味著網(wǎng)絡(luò)在高層次上能提取更豐富的特征而空間尺寸變小意味著特征越來(lái)越全局化。3.2 網(wǎng)絡(luò)結(jié)構(gòu)定義與形狀推導(dǎo)我最終使用的 CNN 結(jié)構(gòu)如下第一層Conv2d(1, 32, kernel_size3, padding1)后接 ReLU再接 MaxPool2d(2)第二層Conv2d(32, 64, kernel_size3, padding1)后接 ReLU再接 MaxPool2d(2)第三層Flatten把二維特征圖拉成一維向量第四層Linear(64×7×7, 128)后接 ReLU第五層Linear(128, 10)輸出十個(gè)類別的分?jǐn)?shù)關(guān)于形狀變化有一個(gè)通用公式可以自己推導(dǎo)。設(shè)輸入特征圖尺寸為 W卷積核大小為 K填充為 P步長(zhǎng)為 S則輸出尺寸為輸出尺寸 (W - K 2P) / S 1當(dāng) W28K3P1S1 時(shí)輸出仍然是 28。經(jīng)過(guò) MaxPool2d(2) 后尺寸減半變成 14。第二次卷積后保持 14再經(jīng)過(guò)池化變成 7。所以 Flatten 前的張量形狀是 [64, 7, 7]64 是通道數(shù)全連接層第一個(gè) Linear 的輸入維度就是 64×7×73136。代碼實(shí)現(xiàn)如下# models/model.py import torch.nn as nn class CNNNet(nn.Module): def __init__(self): super(CNNNet, self).__init__() self.conv_layers nn.Sequential( nn.Conv2d(1, 32, kernel_size3, padding1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size3, padding1), nn.ReLU(), nn.MaxPool2d(2), ) self.fc_layers nn.Sequential( nn.Linear(64 * 7 * 7, 128), nn.ReLU(), nn.Linear(128, 10), ) def forward(self, x): x self.conv_layers(x) x x.view(x.size(0), -1) x self.fc_layers(x) return x代碼里我要特別提醒 view 這一步。x 經(jīng)過(guò)卷積池化后形狀是 [batch_size, 64, 7, 7]view(x.size(0), -1) 表示保持 batch 維度不變把每個(gè)樣本的 64×7×7 展平成 3136 維向量然后才能輸入全連接層。這個(gè)維度不匹配是最常見(jiàn)的報(bào)錯(cuò)大家動(dòng)手寫(xiě)的時(shí)候注意一下。3.3 為什么 CNN 比全連接網(wǎng)絡(luò)更適合圖像很多同學(xué)會(huì)問(wèn)“我也能用多層感知機(jī) MLP 做手寫(xiě)數(shù)字識(shí)別為什么非要用 CNN”確實(shí)MLP 在處理 MNIST 上也能達(dá)到 95% 左右的準(zhǔn)確率但是如果你把輸入圖片稍微平移幾個(gè)像素MLP 的分類結(jié)果可能就變了而 CNN 會(huì)穩(wěn)定很多。原因在于 MLP 把圖像拉平成一維序列后像素之間的空間位置關(guān)系被破壞了。比如一個(gè)像素在第 10 位和第 100 位對(duì)模型來(lái)說(shuō)只是編號(hào)不同模型必須靠大量參數(shù)強(qiáng)行記憶每種數(shù)字模式。CNN 則通過(guò)卷積操作完整保留了二維空間結(jié)構(gòu)。卷積核滑動(dòng)時(shí)相鄰像素的關(guān)聯(lián)被天然建模這種歸納偏置讓模型在小數(shù)據(jù)集上更不容易過(guò)擬合同時(shí)泛化能力更強(qiáng)。所以即使 MNIST 是灰度簡(jiǎn)圖用 CNN 依然是最合理的選擇也能給畢設(shè)的“研究意義”部分提供充足論證素材。4. 訓(xùn)練流程從損失函數(shù)到訓(xùn)練循環(huán)細(xì)節(jié)4.1 損失函數(shù)選擇交叉熵分類問(wèn)題最常用的損失函數(shù)是交叉熵。PyTorch 中可以直接用 nn.CrossEntropyLoss()這個(gè)模塊內(nèi)部把 Softmax 和交叉熵合并在一起。模型輸出的 logits 是一個(gè)長(zhǎng)度 10 的向量每個(gè)位置的數(shù)值代表該類的“未歸一化得分”CrossEntropyLoss 會(huì)先把 logits 通過(guò) Softmax 轉(zhuǎn)成概率分布然后計(jì)算真實(shí)標(biāo)簽分布與預(yù)測(cè)分布的交叉熵。為什么不直接用均方誤差 MSE我在剛開(kāi)始學(xué)習(xí)時(shí)也困惑過(guò)。核心原因是分類問(wèn)題輸出的是離散類別MSE 假設(shè)誤差服從高斯分布適合回歸場(chǎng)景而交叉熵從信息論角度直接衡量?jī)蓚€(gè)概率分布的距離梯度在 Softmax 配合下更有利于分類任務(wù)。換個(gè)直白的說(shuō)法你用交叉熵訓(xùn)練模型預(yù)測(cè)置信度高的樣本會(huì)得到更小的損失預(yù)測(cè)錯(cuò)誤時(shí)的懲罰也更明確訓(xùn)練效率遠(yuǎn)高于 MSE。4.2 優(yōu)化器與學(xué)習(xí)率優(yōu)化器我選擇 Adam這是目前最流行的選擇之一。Adam 相當(dāng)于在 SGD 基礎(chǔ)上加入了一階動(dòng)量和二階動(dòng)量可以在訓(xùn)練中自動(dòng)調(diào)整每個(gè)參數(shù)的學(xué)習(xí)步長(zhǎng)。對(duì) MNIST 這樣的小數(shù)據(jù)集Adam 的默認(rèn)參數(shù)已經(jīng)很好用你不需要過(guò)多糾結(jié)。學(xué)習(xí)率這里我踩過(guò)一次坑。剛開(kāi)始我把學(xué)習(xí)率設(shè)成 0.1結(jié)果 loss 在 2.3 附近原地不動(dòng)甚至偶爾變成 NaN。后來(lái)?yè)Q到 0.001訓(xùn)練在 10 個(gè) epoch 內(nèi)就把測(cè)試準(zhǔn)確率拉到了 99% 左右。如果學(xué)習(xí)率太大會(huì)導(dǎo)致參數(shù)更新跨度過(guò)大越過(guò)最優(yōu)點(diǎn)學(xué)習(xí)率太小則收斂太慢期末周時(shí)間寶貴等不起。我的建議是先從 0.001 起步觀察訓(xùn)練曲線平穩(wěn)下降再在最后幾個(gè) epoch 考慮用 torch.optim.lr_scheduler.StepLR 每若干輪把學(xué)習(xí)率降一半這種操作能讓損失在訓(xùn)練后期進(jìn)一步下降。4.3 完整的訓(xùn)練循環(huán)代碼訓(xùn)練循環(huán)的固定心法我總結(jié)成四步清空梯度、前向傳播、計(jì)算損失、反向傳播和優(yōu)化器步進(jìn)。這幾個(gè)步驟順序不能亂。梯度清空放在最前面如果忘了寫(xiě) zero_grad()PyTorch 默認(rèn)會(huì)累加梯度loss 就會(huì)亂掉。模型要在訓(xùn)練和驗(yàn)證兩種模式間切換通過(guò) model.train() 和 model.eval() 實(shí)現(xiàn)。為什么必須切換因?yàn)?BatchNorm 層和 Dropout 層在訓(xùn)練和測(cè)試時(shí)行為不同。BatchNorm 在訓(xùn)練時(shí)用當(dāng)前 batch 的均值方差在測(cè)試時(shí)用累積的全局統(tǒng)計(jì)量Dropout 在訓(xùn)練時(shí)隨機(jī)丟神經(jīng)元在測(cè)試時(shí)不丟。如果不切換評(píng)估結(jié)果會(huì)有偏差。# utils/train.py 簡(jiǎn)化版核心訓(xùn)練邏輯 def train_one_epoch(model, train_loader, optimizer, criterion, device): model.train() total_loss 0 correct 0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() * images.size(0) pred outputs.argmax(dim1) correct (pred labels).sum().item() return total_loss / len(train_loader.dataset), correct / len(train_loader.dataset)我建議在訓(xùn)練循環(huán)里同時(shí)統(tǒng)計(jì)每個(gè) epoch 的訓(xùn)練準(zhǔn)確率不要只看 lossloss 下降但準(zhǔn)確率不動(dòng)也是可能的。當(dāng)訓(xùn)練準(zhǔn)確率達(dá)到 99% 以上但驗(yàn)證準(zhǔn)確率還在 97% 附近徘徊就說(shuō)明模型開(kāi)始過(guò)擬合訓(xùn)練集了。4.4 訓(xùn)練過(guò)程中的兩個(gè)典型問(wèn)題第一個(gè)問(wèn)題是過(guò)擬合。MNIST 比較友好一般不會(huì)嚴(yán)重過(guò)擬合但如果你把網(wǎng)絡(luò)搞得太寬太深比如每層 512 個(gè)神經(jīng)元就會(huì)開(kāi)始出現(xiàn)訓(xùn)練集 100%、測(cè)試集 98% 這種差距。緩解手段包括增加 Dropout、做數(shù)據(jù)增強(qiáng)、縮小網(wǎng)絡(luò)規(guī)模。我在第七部分會(huì)詳細(xì)展開(kāi)。第二個(gè)問(wèn)題是訓(xùn)練時(shí)間。CPU 上跑我的模型10 個(gè) epoch 大約需要 4 到 6 分鐘完全在可接受范圍內(nèi)。如果你的電腦配置更差建議調(diào)小 batch_size 到 32并減少訓(xùn)練 epoch 到 5先把整個(gè)流程跑通再說(shuō)。千萬(wàn)不要一開(kāi)始就在大參數(shù)上死等調(diào)通流程比追求指標(biāo)重要。5. 評(píng)估與可視化讓模型結(jié)果看得見(jiàn)5.1 準(zhǔn)確率不是唯一指標(biāo)訓(xùn)練結(jié)束后我們要在測(cè)試集上做最終評(píng)估。測(cè)試集是模型從未見(jiàn)過(guò)的 1 萬(wàn)張圖片用它評(píng)估得到的準(zhǔn)確率才是泛化能力的真實(shí)體現(xiàn)。不過(guò)光看準(zhǔn)確率還不夠有說(shuō)服力我在畢設(shè)報(bào)告中補(bǔ)充了精確率、召回率和 F1 值還有混淆矩陣這些指標(biāo)能幫你分析模型到底在哪些類別上犯了錯(cuò)。PyTorch 中可以用 sklearn.metrics 里的 classification_report 和 confusion_matrix 直接計(jì)算。這兩個(gè)函數(shù)很成熟一行代碼就能輸出全部指標(biāo)。不過(guò)我建議你要懂得怎么從混淆矩陣?yán)镒x數(shù)比如第 4 行第 9 列的值是 8意味著有 8 張真實(shí)的數(shù)字 4 被模型誤判成了 9。這種分析寫(xiě)到論文實(shí)驗(yàn)章節(jié)里是很好的素材。5.2 可視化訓(xùn)練曲線和混淆矩陣我習(xí)慣用 matplotlib 畫(huà)兩張圖。第一張是訓(xùn)練集和測(cè)試集的損失值隨 epoch 的變化曲線直觀展示收斂過(guò)程。第二張是混淆矩陣的熱力圖x 軸是預(yù)測(cè)標(biāo)簽y 軸是真實(shí)標(biāo)簽對(duì)角線越亮越好。可視化代碼由于篇幅原因我先不全部貼出真正要掌握的核心就這幾點(diǎn)預(yù)測(cè)結(jié)果是 logits需要用 argmax(dim1) 取概率最大的類別正確率統(tǒng)計(jì)要把預(yù)測(cè)標(biāo)簽和真實(shí)標(biāo)簽逐元素比較混淆矩陣要用測(cè)試集的全部 1 萬(wàn)張圖片來(lái)算不要偷懶只用 1000 張否則誤差會(huì)偏大。我實(shí)際跑出來(lái)的結(jié)果測(cè)試準(zhǔn)確率在 99.0% 到 99.3% 之間在第 2 和第 3 個(gè) epoch 時(shí)準(zhǔn)確率就已經(jīng)能突破 98%后續(xù)訓(xùn)練是穩(wěn)步微調(diào)。5.3 每個(gè) epoch 的輸出效果我給出一次實(shí)際運(yùn)行的參考日志方便大家對(duì)照Epoch訓(xùn)練損失訓(xùn)練準(zhǔn)確率測(cè)試準(zhǔn)確率10.15296.02%97.33%20.04798.68%98.46%30.03299.08%98.82%40.02499.32%98.96%50.01999.43%99.06%60.01699.56%99.10%70.01399.66%99.18%80.01199.73%99.21%90.00999.81%99.25%100.00899.86%99.26%從表格可以看出訓(xùn)練準(zhǔn)確率在持續(xù)提升測(cè)試準(zhǔn)確率也在提升但幅度趨緩。這是正?,F(xiàn)象說(shuō)明模型逐漸收斂。測(cè)試準(zhǔn)確率始終低于訓(xùn)練準(zhǔn)確率這是泛化差距的表現(xiàn)不過(guò)差距很小在可接受范圍內(nèi)。6. 實(shí)驗(yàn)技巧與畢設(shè)答辯延伸6.1 簡(jiǎn)單的數(shù)據(jù)增強(qiáng)改進(jìn)雖然是經(jīng)典數(shù)據(jù)集但實(shí)驗(yàn)部分如果能有一點(diǎn)“改進(jìn)實(shí)驗(yàn)”會(huì)顯得工作量更充裕。一個(gè)常用的手段是數(shù)據(jù)增強(qiáng)也就是對(duì)原始訓(xùn)練圖片做隨機(jī)變換生成更多樣化的訓(xùn)練樣本。對(duì) MNIST 而言比較合適的增強(qiáng)包括隨機(jī)旋轉(zhuǎn) 10 度范圍內(nèi)、隨機(jī)平移兩個(gè)像素、添加少量噪聲。這里我要提醒一個(gè)新手常見(jiàn)誤解手寫(xiě)數(shù)字識(shí)別不應(yīng)該做水平翻轉(zhuǎn)增強(qiáng)。因?yàn)閿?shù)字 6 翻轉(zhuǎn)后會(huì)變成 9數(shù)字 8 翻轉(zhuǎn)后還是 8但數(shù)字 7 翻轉(zhuǎn)后可能變成另一個(gè)數(shù)字。翻轉(zhuǎn)會(huì)破壞類別標(biāo)簽?zāi)P蜁?huì)學(xué)到錯(cuò)誤映射。所以增強(qiáng)策略必須符合任務(wù)本身的語(yǔ)義。PyTorch 中可以在 transforms.Compose 里加上 RandomRotation。代碼上只需改一行但實(shí)驗(yàn)效果可能會(huì)在 99% 的基礎(chǔ)上再穩(wěn)定一點(diǎn)點(diǎn)更重要的是你在論文中可以寫(xiě)“通過(guò)數(shù)據(jù)增強(qiáng)進(jìn)一步提高模型魯棒性”。6.2 模型參數(shù)量計(jì)算老師答辯時(shí)經(jīng)常問(wèn)“你這個(gè)模型有多大、有多少參數(shù)”。你不能只回答一個(gè)模糊的“不大”。參數(shù)量的計(jì)算其實(shí)很簡(jiǎn)單。卷積層參數(shù)量等于卷積核參數(shù)加上偏置計(jì)算公式為參數(shù) 輸入通道 × 輸出通道 × 卷積核高 × 卷積核寬 輸出通道全連接層參數(shù)量等于輸入維度乘輸出維度再加偏置。第一層卷積的參數(shù)量是 1×32×3×332320第二層是 32×64×3×36418496第一個(gè)全連接層是 3136×128128401536輸出層是 128×10101290總參數(shù)約 42 萬(wàn)。這個(gè)規(guī)模非常小存儲(chǔ)模型文件不到 2MB。6.3 答辯常見(jiàn)問(wèn)題與回答思路我整理了一套被高頻提問(wèn)的清單提前準(zhǔn)備總比現(xiàn)場(chǎng)現(xiàn)編強(qiáng)。問(wèn)題建議回答思路為什么選擇 CNN 而不是普通神經(jīng)網(wǎng)絡(luò)圖像有局部相關(guān)性和空間結(jié)構(gòu)CNN 用卷積核提取局部特征參數(shù)共享減少參數(shù)量池化增強(qiáng)平移不變性卷積層和池化層分別有什么作用卷積負(fù)責(zé)特征提取池化負(fù)責(zé)降維和保留重要特征兩者配合減少計(jì)算量并增強(qiáng)泛化為什么使用 ReLU 激活函數(shù)計(jì)算簡(jiǎn)單、能緩解梯度消失相比 sigmoid/tanh 收斂更快訓(xùn)練中過(guò)擬合怎么解決降低模型復(fù)雜度、加入 Dropout、數(shù)據(jù)增強(qiáng)、早停、增加正則化測(cè)試集和驗(yàn)證集有什么不同驗(yàn)證集用于訓(xùn)練過(guò)程中調(diào)參選模型測(cè)試集只用于最終評(píng)估絕不參與訓(xùn)練這些問(wèn)題沒(méi)有標(biāo)準(zhǔn)答案但思路對(duì)了就能拿分。你在平時(shí)訓(xùn)練時(shí)多記錄幾組實(shí)驗(yàn)對(duì)比比如不同學(xué)習(xí)率下的收斂情況答辯時(shí)能拿出來(lái)展示說(shuō)服力遠(yuǎn)勝于口頭描述。7. 典型問(wèn)題排查與項(xiàng)目擴(kuò)展方向7.1 我在實(shí)際開(kāi)發(fā)中遇到的坑這個(gè)項(xiàng)目看著簡(jiǎn)單真動(dòng)手時(shí)還是會(huì)遇到各種意外。我先說(shuō)最常見(jiàn)的。MNIST 數(shù)據(jù)集默認(rèn)從網(wǎng)上下載如果網(wǎng)絡(luò)不穩(wěn)定下載到一半中斷torchvision 會(huì)報(bào)錯(cuò)或者留下殘缺文件。解決辦法是刪除 data/MNIST 目錄下的不完整文件重新下載。如果實(shí)在沒(méi)有網(wǎng)絡(luò)環(huán)境可以從有網(wǎng)環(huán)境拿到完整的 MNIST 文件然后手動(dòng)放到正確目錄確保文件名稱和結(jié)構(gòu)一致。第二個(gè)容易踩的坑是設(shè)備問(wèn)題。默認(rèn)訓(xùn)練跑在 CPU 上有些同學(xué)的電腦內(nèi)存只有 8GBbatch_size 設(shè)得過(guò)大可能導(dǎo)致內(nèi)存溢出。我的建議是先用 batch_size32 做一次冒煙測(cè)試確保整個(gè)流程能跑通再?zèng)Q定要不要加大。第三個(gè)坑是 loss 出現(xiàn) NaN。這個(gè)大多數(shù)時(shí)候是學(xué)習(xí)率過(guò)大或者數(shù)據(jù)沒(méi)歸一化導(dǎo)致的。如果 Pixel 值還是 0~255loss 直接 NaN 的概率很高。我還遇到過(guò)在 Jupyter 中運(yùn)行多次訓(xùn)練代碼模型參數(shù)和優(yōu)化器狀態(tài)累積導(dǎo)致結(jié)果一次比一次奇怪。這種狀態(tài)污染類問(wèn)題建議每次訓(xùn)完重新實(shí)例化模型不要反復(fù)用同一對(duì)象接著訓(xùn)練。7.2 用 PyTorch 快速手寫(xiě)識(shí)別模型改進(jìn)方向如果你想讓這個(gè)項(xiàng)目從課程作業(yè)晉升為畢業(yè)設(shè)計(jì)亮點(diǎn)可以在幾個(gè)方向上擴(kuò)展。第一是交互界面用 Tkinter 或 PyQt 畫(huà)一個(gè)手寫(xiě)板鼠標(biāo)在上面寫(xiě)數(shù)字模型實(shí)時(shí)識(shí)別結(jié)果。這種 Demo 在答辯現(xiàn)場(chǎng)效果非常好我見(jiàn)過(guò)很多同學(xué)靠這一點(diǎn)把分?jǐn)?shù)拉高。第二是 Web 部署用 Flask 或 FastAPI 封裝模型接口瀏覽器上傳圖片返回識(shí)別結(jié)果這部分能體現(xiàn)工程化能力。第三是模型結(jié)構(gòu)改進(jìn)可以嘗試 ResNet 風(fēng)格的殘差連接或者把普通卷積替換為深度可分離卷積從而在參數(shù)量基本不變的情況下提升精度。這個(gè)項(xiàng)目的設(shè)計(jì)思路同樣適用于其他圖像分類任務(wù)。比如把 MNIST 換成 Fashion-MNIST你就需要把輸入通道、類別數(shù)保持一致只是模型要學(xué)會(huì)區(qū)分不同衣著類別換到無(wú)人機(jī)航拍數(shù)據(jù)集、道路裂縫數(shù)據(jù)集等場(chǎng)景時(shí)重點(diǎn)也不再是網(wǎng)絡(luò)結(jié)構(gòu)本身而是數(shù)據(jù)標(biāo)注質(zhì)量和輸入圖片的預(yù)處理方式。所以說(shuō)到底你在這個(gè)項(xiàng)目里建立的數(shù)據(jù)處理、模型訓(xùn)練、評(píng)估分析閉環(huán)才是真正能遷移的能力。7.3 關(guān)于源碼、注釋和數(shù)據(jù)集交付的補(bǔ)充建議最后說(shuō)一個(gè)很容易被忽略的“非技術(shù)”問(wèn)題交付物形式。老師或者評(píng)審最終看到的不只是代碼能不能跑還包括代碼注釋清不清楚、數(shù)據(jù)集是否完整、README 是否看得懂。我給自己的每個(gè)函數(shù)都寫(xiě)了 docstring關(guān)鍵訓(xùn)練步驟也加了中文注釋這不是制造工作量而是為了讓自己兩個(gè)星期后回頭看代碼還能一眼看懂當(dāng)初的意圖。數(shù)據(jù)集方面默認(rèn)情況下 torchvision 會(huì)自動(dòng)下載但我額外把數(shù)據(jù)目錄單獨(dú)整理好并在 README 里寫(xiě)了“目錄結(jié)構(gòu)說(shuō)明”。如果你是離線交付建議把 MNIST 數(shù)據(jù)集一起打包避免對(duì)方運(yùn)行時(shí)才去下載。這里補(bǔ)充一點(diǎn)個(gè)人的經(jīng)驗(yàn)項(xiàng)目源碼包如果超過(guò) 200MB建議分卷壓縮評(píng)分老師用微信或郵件收件時(shí)不會(huì)被單文件大小卡住。我自己在完成這個(gè)項(xiàng)目時(shí)最大的感受是不要被“深度學(xué)習(xí)”四個(gè)字嚇住。借助 PyTorch 和標(biāo)準(zhǔn)數(shù)據(jù)集哪怕只有基礎(chǔ) Python 知識(shí)也能在不到一周的時(shí)間內(nèi)實(shí)現(xiàn)一個(gè)表現(xiàn)很好的圖像分類系統(tǒng)。第一次跑通訓(xùn)練循環(huán)看到準(zhǔn)確率攀升的那一刻你會(huì)發(fā)現(xiàn)之前踩過(guò)的所有坑都值得。希望這篇內(nèi)容能幫你少走幾步彎路也別為了保證安全而錯(cuò)過(guò)大作業(yè)的學(xué)分。本文還有配套的精品資源點(diǎn)擊獲取