字識別工程:MNIST與PyTorch實戰(zhàn)全解析)
簡介圖像分類是機(jī)器學(xué)習(xí)中最基礎(chǔ)也最具代表性的任務(wù)之一手寫數(shù)字識別作為入門經(jīng)典能夠直觀展示從數(shù)據(jù)預(yù)處理、模型構(gòu)建到訓(xùn)練推理的完整流程。MNIST數(shù)據(jù)集是這一領(lǐng)域的基準(zhǔn)其28x28灰度圖像和標(biāo)準(zhǔn)的訓(xùn)練測試劃分讓研究者可以快速驗證算法效果。本文以PyTorch工程實踐為主線介紹如何將原始二進(jìn)制數(shù)據(jù)解析、歸一化處理、全連接網(wǎng)絡(luò)與卷積神經(jīng)網(wǎng)絡(luò)的選型對比、訓(xùn)練調(diào)參與模型部署等環(huán)節(jié)串聯(lián)成一套可復(fù)用的工程框架。此外通過分析驗證集loss曲線與準(zhǔn)確率波動讀者能理解過擬合與欠擬合的典型形態(tài)并掌握模型保存、單張圖片推理和打包發(fā)布的工程化技巧。在此基礎(chǔ)上這套方案可輕松遷移到Fashion-MNIST或更復(fù)雜的圖像分類任務(wù)從而理解深度學(xué)習(xí)工程的通用范式。1. 拿到工程文件后先搞明白這套代碼到底在做什么先說個我經(jīng)常在帶新人時遇到的場景很多人下了一堆“手寫數(shù)字識別”的代碼解壓之后直接雙擊train.py看到屏幕上滾出幾個epoch、打出一行accuracy就覺得“跑通了”。但你要是讓他說清楚這套工程的文件結(jié)構(gòu)為什么這么拆、模型輸入為什么是784維、訓(xùn)練集為什么要除以255他大概率答不上來。這套Python手寫數(shù)字識別項目本質(zhì)上是一套完整的圖像分類工程。它的業(yè)務(wù)目標(biāo)很樸素給定一張包含手寫數(shù)字的圖片讓程序判斷它到底是0到9中的哪一個。但“工程化”這三個字意味著代碼不只是能跑通而是要覆蓋從數(shù)據(jù)預(yù)處理、模型構(gòu)建、訓(xùn)練驗證到推理部署的全鏈路并且每一條路徑都有清晰的輸入輸出約定。整個工程的核心鏈路可以拆成四段數(shù)據(jù)管線手寫數(shù)字圖片 - Numpy數(shù)組 - 歸一化 - 張量模型主體一個接收784維輸入、輸出10類概率的分類器訓(xùn)練引擎通過交叉熵?fù)p失和梯度下降不斷修正權(quán)重推理服務(wù)加載訓(xùn)練好的權(quán)重對新圖片做預(yù)測并輸出可視化結(jié)果。我見過很多初學(xué)者的誤區(qū)是把“訓(xùn)練”和“工程”畫等號其實訓(xùn)練只是其中一環(huán)。真正決定這套代碼能不能被別人復(fù)現(xiàn)、能不能遷移到別的任務(wù)上取決于文件組織是否合理、配置項是否獨立、數(shù)據(jù)路徑是否可配置。判斷一套工程文件優(yōu)劣最簡單的方法把你電腦上的絕對路徑全部換成相對路徑看代碼還能不能跑起來。能跑說明工程底子合格不能跑說明這只是個腳本合集不是工程。2. MNIST數(shù)據(jù)集的獲取與預(yù)處理實操別讓數(shù)據(jù)拖了后腿手寫數(shù)字識別最經(jīng)典的數(shù)據(jù)集就是MNIST。它包含60000張訓(xùn)練圖片和10000張測試圖片每張圖片是28x28像素的灰度圖。這里有個很多教程沒強(qiáng)調(diào)的細(xì)節(jié)MNIST原始文件不是圖片格式而是特定的二進(jìn)制文件格式。你下載下來會看到四個文件分別是train-images-idx3-ubyte.gz、train-labels-idx1-ubyte.gz、t10k-images-idx3-ubyte.gz和t10k-labels-idx1-ubyte.gz雖然網(wǎng)友整理版可能已經(jīng)幫你解壓并轉(zhuǎn)換成了圖片但工程中更推薦直接處理原始二進(jìn)制。2.1 為什么工程里要保留原始二進(jìn)制文件的解析邏輯因為二進(jìn)制的讀取速度遠(yuǎn)遠(yuǎn)快于逐張讀取圖片文件。圖片格式意味著系統(tǒng)要調(diào)用圖像解碼庫把JPEG或PNG的數(shù)據(jù)解碼成像素矩陣這個I/O開銷在批量訓(xùn)練時是很可觀的。而二進(jìn)制文件本身已經(jīng)是按固定字節(jié)結(jié)構(gòu)排列的像素值你只需要按偏移量切片再用Numpy的frombuffer轉(zhuǎn)成數(shù)組速度會快一個量級。工程文件里通常會在data_loader.py中封裝這樣一個函數(shù)import numpy as np import struct def load_mnist_images(filepath): with open(filepath, rb) as f: magic, num, rows, cols struct.unpack(IIII, f.read(16)) data np.frombuffer(f.read(), dtypenp.uint8).reshape(num, rows * cols) return data def load_mnist_labels(filepath): with open(filepath, rb) as f: magic, num struct.unpack(II, f.read(8)) labels np.frombuffer(f.read(), dtypenp.uint8) return labels這里有個比較隱蔽的坑文件頭部的magic number是按大端序存儲的所以解包必須用IIII。我見過不少人的代碼在這里用IIII解包結(jié)果第一個數(shù)字讀出來是個異常值整個數(shù)據(jù)集的shape都亂了。還有一點labels文件只有兩個頭部字段不像images有rows和cols多讀一個字段就會導(dǎo)致buf偏移錯誤。2.2 像素歸一化到底在做什么MNIST原始像素值的范圍是0到255。如果直接喂給神經(jīng)網(wǎng)絡(luò)有兩個問題數(shù)值量級太大會讓初始化權(quán)重對應(yīng)的梯度更新變得不穩(wěn)定不同維度的輸入范圍不一致會讓模型收斂變慢。所以工程里幾乎無例外都會做歸一化把像素值壓到0到1之間。做法很簡單X_train X_train.astype(np.float32) / 255.0 X_test X_test.astype(np.float32) / 255.0如果你用PyTorch還需要再做一步把Numpy數(shù)組轉(zhuǎn)成Tensor并且把標(biāo)簽也轉(zhuǎn)成LongTensor。這里有個和后續(xù)模型匹配的概念必須講清楚標(biāo)簽是0到9的標(biāo)量不是one-hot向量。模型最后一層輸出的是10個類別的logitsPyTorch的CrossEntropyLoss函數(shù)會內(nèi)部幫你組合LogSoftmax和NLLLoss所以直接喂標(biāo)簽索引就行如果你自己把標(biāo)簽轉(zhuǎn)成one-hot向量再和Softmax輸出算損失就要自己實現(xiàn)對應(yīng)的損失函數(shù)容易出錯。2.3 數(shù)據(jù)維度怎么確認(rèn)接數(shù)據(jù)的時候最好打印一次數(shù)據(jù)的shape和dtype不要憑記憶。我調(diào)試過不少次問題最后都出在某個環(huán)節(jié)維度對不上圖片load出來是(60000, 784)標(biāo)簽是(60000,)網(wǎng)絡(luò)前向傳播需要的輸入是(batch_size, 784)如果batch_size為32那么一個batch的tensor shape就是(32, 784)。如果維度對不上會直接報矩陣乘法錯誤。工程文件里一般會在主訓(xùn)練腳本開頭加一行斷言assert X_train.shape[0] y_train.shape[0] assert X_train.shape[1] 28 * 283. 模型選型與訓(xùn)練細(xì)節(jié)從全連接網(wǎng)絡(luò)到卷積網(wǎng)絡(luò)的實測差異很多手寫數(shù)字識別工程會從多層感知機(jī)開始。這個選擇是有道理的數(shù)字識別是入門任務(wù)用全連接網(wǎng)絡(luò)可以清晰理解網(wǎng)絡(luò)的前向過程、反向傳播和參數(shù)更新不會一開始就被卷積、池化等概念淹沒。但如果你想在MNIST上拿到比較好看的準(zhǔn)確率純?nèi)B接網(wǎng)絡(luò)和簡單CNN的差距還是存在的。3.1 多層感知機(jī)怎么設(shè)計最穩(wěn)一個典型的MLP結(jié)構(gòu)可以設(shè)計成三層輸入層784個神經(jīng)元、隱藏層128個神經(jīng)元、輸出層10個神經(jīng)元。中間加ReLU激活函數(shù)和Dropout正則化。在PyTorch里寫出來是這樣import torch.nn as nn class MLP(nn.Module): def __init__(self): super().__init__() self.fc1 nn.Linear(784, 128) self.relu nn.ReLU() self.dropout nn.Dropout(0.2) self.fc2 nn.Linear(128, 10) def forward(self, x): x x.view(x.size(0), -1) x self.fc1(x) x self.relu(x) x self.dropout(x) x self.fc2(x) return x很多人拿到工程文件后最關(guān)心的一個問題是為什么第一層要view因為訓(xùn)練時輸入進(jìn)來的tensor形狀可能是(batch_size, 1, 28, 28)如果直接丟給Linear層會報錯必須把它壓平成(batch_size, 784)。這個view操作其實就是把28x28的矩陣?yán)梢粭l784維的向量。MLP在MNIST上做到97%左右的準(zhǔn)確率沒有問題我當(dāng)時實測大概在97.2%。但再往上就比較費勁了因為它丟失了圖像的空間結(jié)構(gòu)信息每個像素位置都是獨立特征無法捕捉相鄰像素之間的空間相關(guān)性。3.2 什么時候該上CNN如果你想在MNIST上沖擊99%以上的準(zhǔn)確率就得換CNN。一個經(jīng)典的LeNet-5結(jié)構(gòu)可以很好地完成任務(wù)。它的核心思想是用卷積核在圖像上滑動提取局部特征。對于28x28的MNIST圖片第一層卷積可以輸出多個特征圖每個特征圖捕捉一種模式比如橫線、豎線、圓角等。我自己的工程里用的是LeNet-5的變體class LeNet5(nn.Module): def __init__(self): super().__init__() self.conv1 nn.Conv2d(1, 6, kernel_size5, padding2) self.pool1 nn.MaxPool2d(2) self.conv2 nn.Conv2d(6, 16, kernel_size5) self.pool2 nn.MaxPool2d(2) self.fc1 nn.Linear(16 * 5 * 5, 120) self.fc2 nn.Linear(120, 84) self.fc3 nn.Linear(84, 10) def forward(self, x): x self.pool1(torch.relu(self.conv1(x))) x self.pool2(torch.relu(self.conv2(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注意這里x.view(x.size(0), -1)是在全連接層之前展平特征圖。初始輸入是(1, 28, 28)經(jīng)過一次卷積和池化變成(6, 14, 14)經(jīng)過第二次卷積和池化變成(16, 5, 5)展平后是1655400維。后接120、84、10的三層全連接。這個結(jié)構(gòu)在MNIST上輕輕松松就能達(dá)到99%以上。CNN和MLP的差異用一句話概括MLP是把整個圖像揉成一團(tuán)丟掉空間關(guān)系CNN是通過滑動窗口保留“哪里有什么形狀”的信息。對書寫數(shù)字這種高度依賴形態(tài)特征的任務(wù)CNN優(yōu)勢非常明顯。3.3 訓(xùn)練輪次、損失函數(shù)、優(yōu)化器怎么配超參數(shù)配置在工程里通常是單獨抽出來的。為什么單獨抽因為你不會只想跑一次你需要反復(fù)調(diào)整把配置集中到一個文件或一段常量區(qū)域調(diào)整成本會低很多。我常用的配置是這樣的參數(shù)取值說明batch_size64顯存占用適中梯度的隨機(jī)性合適learning_rate0.001Adam優(yōu)化器下比較穩(wěn)epochs10數(shù)據(jù)集不大10輪足夠收斂optimizerAdam自帶動量收斂速度快loss_fnCrossEntropyLoss多分類標(biāo)準(zhǔn)損失訓(xùn)練循環(huán)需要注意的細(xì)節(jié)是每一輪epoch結(jié)束時要區(qū)分訓(xùn)練集和驗證集的loss。不能只看訓(xùn)練集準(zhǔn)確率因為模型可能已經(jīng)過擬合了。工程里一般會在每個epoch后把模型切到eval模式關(guān)閉dropout用驗證集算一次準(zhǔn)確率然后存下best_model。這里有個用PyTorch的常見細(xì)節(jié)訓(xùn)練時要調(diào)用model.train()eval時要調(diào)用model.eval()否則Dropout和BatchNorm的行為會不一致導(dǎo)致結(jié)果偏高。另外學(xué)習(xí)率衰減也很重要。前期用0.001快速收斂后期降到0.0001精細(xì)微調(diào)能再拉高一點準(zhǔn)確率。實現(xiàn)上用PyTorch的torch.optim.lr_scheduler.StepLR每3個epoch乘以0.1即可。4. 訓(xùn)練過程的完整調(diào)參與評估從loss曲線到準(zhǔn)確率波動的解讀訓(xùn)練代碼能跑不代表訓(xùn)練過程健康這是很多初學(xué)者的一個認(rèn)知死角。工程文件里通常提供訓(xùn)練過程的可視化腳本輸出loss曲線和準(zhǔn)確率曲線但更重要的是你要會看這些曲線理解梯度和過擬合的跡象。4.1 loss曲線的三種典型形態(tài)訓(xùn)練過程結(jié)束之后我們把每個epoch的loss值和驗證準(zhǔn)確率畫出來。這里我總結(jié)三種常見的曲線形態(tài)你在自己的訓(xùn)練中也會碰到相同的模式理想形態(tài)訓(xùn)練loss和驗證loss同步下降最后都收斂到較低水平驗證準(zhǔn)確率穩(wěn)定在99%上下。這說明模型容量、數(shù)據(jù)量、學(xué)習(xí)率三者匹配得很好不需要做額外調(diào)整。過擬合形態(tài)訓(xùn)練loss持續(xù)下降但驗證loss下降到某個點后開始反彈。這個轉(zhuǎn)折點提示你模型開始“記”訓(xùn)練數(shù)據(jù)而不是“學(xué)”規(guī)律。應(yīng)對方案是增加Dropout強(qiáng)度、增加數(shù)據(jù)增強(qiáng)或者減少隱藏層神經(jīng)元數(shù)量。欠擬合形態(tài)訓(xùn)練loss和驗證loss都居高不下驗證準(zhǔn)確率一直在97%以下徘徊。這說明模型容量不夠或者學(xué)習(xí)率太小、收斂太慢。此時優(yōu)先增加網(wǎng)絡(luò)層數(shù)或每層的神經(jīng)元數(shù)量。4.2 為什么驗證集準(zhǔn)確率比訓(xùn)練集重要工程里看模型好壞標(biāo)準(zhǔn)不是訓(xùn)練集上的表現(xiàn)而是驗證集上的表現(xiàn)因為模型未來遇到的是沒有見過的數(shù)據(jù)。MNIST數(shù)據(jù)集本身已經(jīng)劃分好了train和test但很多工程還會再從train中切一個validation出來。如果你不想額外切直接用test集做驗證也是可以的但嚴(yán)格來說測試集應(yīng)該只用于最終評估不能進(jìn)訓(xùn)練循環(huán)否則你在根據(jù)測試結(jié)果調(diào)參的過程中其實已經(jīng)發(fā)生了信息泄漏。我當(dāng)時在自己的工程里是按照6:1的比例從訓(xùn)練集中切分驗證集保留的10000條數(shù)據(jù)作為測試集。這樣每輪epoch都能直觀看到驗證準(zhǔn)確率最后再用測試集跑一次得到的是模型真實的泛化能力。4.3 訓(xùn)練過程中的穩(wěn)定性和收斂性觀測除了準(zhǔn)確率還要看一下訓(xùn)練過程中的數(shù)值穩(wěn)定性。比如loss如果出現(xiàn)NaN基本是學(xué)習(xí)率過大或者數(shù)據(jù)預(yù)處理出了問題要立即停止排查。數(shù)值穩(wěn)定性的另一個常見問題是梯度爆炸或梯度消失全連接網(wǎng)絡(luò)在層數(shù)較深時更容易出現(xiàn)但在MNIST這種淺層模型中比較少見到。我在工程里還加了一行邏輯在測試集上評估準(zhǔn)確率時最好設(shè)置一個閾值比如0.98如果低于這個值則打印警告。這個不是給機(jī)器看的是給人看的提醒你是不是該調(diào)參了。工程化的意義就在于此它不替你判斷但它把判斷依據(jù)信息以清晰方式暴露給你。5. 模型保存與單張圖片推理的工程化處理訓(xùn)練完成只是上半場模型要能給別人用必須解決兩個問題權(quán)重文件怎么存、別人拿一張新圖片怎么預(yù)測。很多工程文件在這里的代碼比較亂我重點說一下合理的做法。5.1 保存PyTorch模型時別只保存state_dictPyTorch有幾種保存方式常見的是torch.save(model.state_dict(), mnist_cnn.pt)只保存權(quán)重torch.save(model, mnist_cnn.pth)保存整個模型結(jié)構(gòu)加權(quán)重onnx.export(model, dummy_input, mnist_cnn.onnx)導(dǎo)出成跨框架的ONNX格式。我強(qiáng)烈建議在工程里使用state_dict因為它和模型結(jié)構(gòu)解耦加載時必須先創(chuàng)建相同結(jié)構(gòu)的模型實例再load。雖然比直接保存整個模型多一步但它在版本遷移、結(jié)構(gòu)修改的時候更靈活。你在每個最佳epoch保存best_model.pt之外最好同時保存一份final_model.pt以免中途訓(xùn)練中斷丟了最佳結(jié)果。5.2 單張手寫數(shù)字圖片的預(yù)處理流程推理階段最容易翻車的點不是模型代碼而是圖片預(yù)處理。用戶傳過來的圖片不可能是標(biāo)準(zhǔn)的MNIST格式它可能是手機(jī)拍的、用畫圖工具畫的、或者從PDF截圖的。所以工程里推理部分的預(yù)處理器必須做下面這幾件事順序也不可隨意調(diào)換讀取圖片轉(zhuǎn)為灰度圖反色處理如果背景是白色、筆跡是黑色但MNIST是黑底白字需要顛倒縮放到28x28二值化或保持灰度值歸一化到0到1加一個batch維度(1, 1, 28, 28)。這里最容易被忽略的是反色。我剛開始做推理Demo的時候用畫圖工具寫了個“7”預(yù)測出來是“1”排查半天發(fā)現(xiàn)白色背景255直接變成了高亮值模型看到的是“白字黑底”輸入分布完全顛倒。加一步cv2.bitwise_not()或者在歸一化時用1 - img/255.0就能解決。這屬于那種不踩一次坑就不知道的細(xì)節(jié)。完整的推理代碼大致長這樣import cv2 import torch import numpy as np def preprocess_image(image_path): img cv2.imread(image_path, cv2.IMREAD_GRAYSCALE) img cv2.bitwise_not(img) # 反色 img cv2.resize(img, (28, 28), interpolationcv2.INTER_AREA) img img.astype(np.float32) / 255.0 img torch.from_numpy(img).unsqueeze(0).unsqueeze(0) return img model LeNet5() model.load_state_dict(torch.load(mnist_cnn.pt, map_locationcpu)) model.eval() with torch.no_grad(): img preprocess_image(test_7.png) output model(img) pred torch.argmax(output, dim1).item() print(f預(yù)測結(jié)果: {pred})訓(xùn)練時的model.eval()同樣適用在推理階段。這里如果漏了torch.no_grad()模型還是會正常給出結(jié)果但會記錄梯度圖白白消耗內(nèi)存推理時間也會變長。5.3 用OpenCV畫圖板實時測試模型圖片文件推理只是工程的一部分實際應(yīng)用里還有實時輸入的需求。我當(dāng)時又做了一層簡單的GUI用OpenCV創(chuàng)建一個窗口鼠標(biāo)按住畫數(shù)字松開后按Enter鍵進(jìn)行識別結(jié)果實時顯示在窗口標(biāo)題上。雖然不能和TensorFlow的Playground對比但代碼量很少、依賴很少非常適合作為工程演示的一部分。實現(xiàn)思路也不復(fù)雜先標(biāo)記鼠標(biāo)按下時在Canvas上畫圓結(jié)束后把Canvas區(qū)域作為輸入圖片傳給同一套預(yù)處理流程。這個改進(jìn)讓你不用每次都準(zhǔn)備圖片文件調(diào)試手感提升非常明顯。6. 工程文件的目錄組織、依賴管理與打包發(fā)布既然標(biāo)題寫的是“工程文件”那這章必須認(rèn)真講。一個合格的手寫數(shù)字識別工程目錄不能是一堆.py文件堆在根目錄。我推薦的結(jié)構(gòu)是這樣的mnist_project/ ├── checkpoints/ # 保存訓(xùn)練好的模型權(quán)重 ├── data/ # MNIST原始數(shù)據(jù)或下載腳本 ├── models/ # 網(wǎng)絡(luò)結(jié)構(gòu)定義 │ └── lenet5.py ├── utils/ # 數(shù)據(jù)處理、可視化工具 │ ├── data_loader.py │ └── visualizer.py ├── config.py # 超參數(shù)集中管理 ├── train.py # 訓(xùn)練入口 ├── predict.py # 單張圖片推理入口 ├── requirements.txt └── README.md這套結(jié)構(gòu)的好處是職責(zé)清晰網(wǎng)絡(luò)結(jié)構(gòu)、數(shù)據(jù)處理、訓(xùn)練流程、推理流程各自獨立換網(wǎng)絡(luò)結(jié)構(gòu)時不用動數(shù)據(jù)代碼換數(shù)據(jù)時不用動模型代碼。很多教程代碼喜歡把所有函數(shù)都放進(jìn)一個文件跑通是快但后續(xù)擴(kuò)展和維護(hù)的代價很大。如果它是一個給別人下載的工程那更要注意這一點。requirements.txt的內(nèi)容至少要包含torch numpy opencv-python matplotlib這幾樣是缺一不可的。建議在文件里固定版本號避免不同用戶環(huán)境差異導(dǎo)致的問題。我自己一般會寫torch2.0,2.3這樣的范圍既兼容新版又不會因為某個大版本API變化直接報錯。關(guān)于打包發(fā)布如果你想讓沒有Python環(huán)境的用戶也能直接運行可以嘗試用PyInstaller把inference腳本打包成exe。這里有個和資源路徑有關(guān)的坑PyTorch的模型文件在打包時不會自動包含進(jìn)去需要在spec文件里把checkpoint作為data文件加進(jìn)去運行時通過sys._MEIPASS獲取臨時解壓路徑。如果忘記這一步別人雙擊exe時會報“文件不存在”的錯誤。打包命令大概是這樣pyinstaller -F predict.py --add-data checkpoints/mnist_cnn.pt;checkpoints --hidden-importtorch --hidden-importcv2注意Windows下--add-data的文件分隔符是分號Linux和macOS是冒號。這個細(xì)節(jié)卡了我差不多一個下午你不遇到真的不會想到。7. 推理結(jié)果的可視化與交互讓工程看起來更完整一套工程如果只有命令行輸出總感覺差點意思。當(dāng)時我把可視化部分補(bǔ)上之后整個項目的完整度明顯提升了。用Matplotlib把待預(yù)測圖片顯示出來同時把10個類別的預(yù)測概率用條形圖展示能直觀看出模型對某個數(shù)字的置信度。特別是當(dāng)模型預(yù)測錯誤時看概率分布能立刻定位問題。我在工程里實現(xiàn)了一個predict_multiple.py腳本支持傳入一個文件夾批量識別所有外部圖片并生成一張匯總圖。匯總圖左邊是原始圖片右邊是預(yù)測概率分布如果某個數(shù)字的置信度低于70%就把預(yù)測結(jié)果標(biāo)紅。這個做法在文檔演示和教學(xué)場景里都很有用能直觀體現(xiàn)模型的可靠邊界。如果后續(xù)想更進(jìn)一步可以用Flask做一個簡單的Web服務(wù)。前端頁面上放一個Canvas鼠標(biāo)手寫數(shù)字點擊識別按鈕后通過POST請求把圖片base64編碼發(fā)給后端后端把圖片解碼、預(yù)處理、推理返回預(yù)測結(jié)果。核心服務(wù)代碼和本地推理幾乎一致只是加了一層HTTP封裝。這樣做的好處是演示的時候不用裝Python環(huán)境打開瀏覽器就能用。不過在工程里引入Web層時要留意請求體大小限制和并發(fā)處理。手寫數(shù)字圖片很小一般不會出問題但如果你把這個架構(gòu)遷移到更大圖片的分類任務(wù)上就需要在服務(wù)端做圖片壓縮和隊列化處理了。8. 實測踩坑記錄數(shù)據(jù)、訓(xùn)練、打包三層里的常見問題寫到最后把我在整個工程實施過程中遇到過的幾個真實問題整理一下希望能幫你少走彎路。8.1 數(shù)據(jù)集加載的坑訓(xùn)練時如果發(fā)現(xiàn)loss完全不下降第一件事檢查數(shù)據(jù)有沒有喂對。我之前遇到過一次圖片讀進(jìn)來之后忘記除以255網(wǎng)絡(luò)訓(xùn)練前幾步loss在2.3左右訓(xùn)到最后只降到1.2準(zhǔn)確率卡在85%。就是因為像素值范圍不對梯度方向被大數(shù)值主導(dǎo)了收斂極慢。還有一種更隱蔽的情況是標(biāo)簽和圖片錯位通常發(fā)生在你手動從網(wǎng)上找數(shù)據(jù)集、目錄文件名和標(biāo)簽映射錯的時候。判斷方法是打印前20張圖片的標(biāo)簽并同時輸出數(shù)組第一個像素的平均值肉眼對應(yīng)一下。8.2 推理結(jié)果不準(zhǔn)的坑模型在測試集上準(zhǔn)確率99%但識別自己手寫的數(shù)字卻總出錯。這大概率是預(yù)處理不夠規(guī)范。手寫輸入和MNIST原始訓(xùn)練集的差異包括字體粗細(xì)、位置偏移、筆畫噪聲。其中位置偏移影響最大MNIST訓(xùn)練集里數(shù)字是居中顯示的如果畫圖時數(shù)字偏上或偏下識別準(zhǔn)確率就會下降。緩解辦法之一是在預(yù)處理時做一次質(zhì)心平移計算出前景像素的均值坐標(biāo)把質(zhì)心移到圖像中心。代碼非常簡單coords cv2.findNonZero(img) x, y, w, h cv2.boundingRect(coords) img img[y:yh, x:xw] img cv2.resize(img, (20, 20)) canvas np.zeros((28, 28), dtypenp.uint8) canvas[4:24, 4:24] img這個操作的本質(zhì)是模擬MNIST的預(yù)處理方式。加了這個步驟后識別率會明顯提升。8.3 打包exe的坑除了前面提到的--add-data路徑分隔符問題還有一個常見坑是打包出來的exe體積異常大動輒幾百MB。這是因為PyTorch和OpenCV依賴庫本身體積很大PyInstaller默認(rèn)把它們?nèi)看蜻M(jìn)去。如果只是給內(nèi)部演示用其實無所謂如果真的很在意體積可以考慮用ONNX Runtime替代PyTorch做推理把模型導(dǎo)出成ONNX格式這樣依賴庫會小很多。我把模型用torch.onnx.export導(dǎo)出后用onnxruntime推理打包體積從440MB降到了80MB左右。8.4 隨機(jī)種子固定問題工程復(fù)現(xiàn)的另一個隱藏要求是固定隨機(jī)種子。如果不加torch.manual_seed(42)每次運行結(jié)果會有細(xì)微差異雖然準(zhǔn)確率都差不多但別人復(fù)現(xiàn)時看到的曲線可能不一致容易被誤以為是代碼bug。在train.py開頭固定種子是個好習(xí)慣import random import numpy as np import torch random.seed(42) np.random.seed(42) torch.manual_seed(42)9. 從數(shù)字識別到其它分類任務(wù)這套工程能怎么擴(kuò)展手寫數(shù)字識別是一個基準(zhǔn)項目但它完全可以作為模板擴(kuò)展到其他分類任務(wù)這也是這套工程文件真正的延伸價值。最簡單的擴(kuò)展是換數(shù)據(jù)集。比如把MNIST的數(shù)據(jù)加載換成Fashion-MNIST只需要調(diào)整類別名模型結(jié)構(gòu)基本不用變就能識別衣服、鞋子、包等10類物品。因為Fashion-MNIST的圖片尺寸和通道數(shù)和MNIST完全一致。這個遷移成本非常低非常適合驗證你的工程結(jié)構(gòu)是否足夠通用。如果想識別中文字符問題會復(fù)雜一些。中文字符類別數(shù)多動輒上千類且筆畫結(jié)構(gòu)復(fù)雜28x28的分辨率可能不夠需要把輸入尺寸擴(kuò)大到64x64或者更大同時模型也要加深。此時卷積層的kernel size、池化層的步長都可能需要調(diào)整。但整體工程的骨架依然是通用的你只需要改數(shù)據(jù)管線和模型結(jié)構(gòu)訓(xùn)練流程、驗證邏輯、推理框架都能復(fù)用。更進(jìn)一步如果輸入不是灰度圖而是彩色圖片比如識別水果種類就需要在第一層卷積前把輸入通道從1改成3同時數(shù)據(jù)預(yù)處理階段要保留RGB三個通道。這個改動也不復(fù)雜但要注意歸一化方式RGB圖像的均值和標(biāo)準(zhǔn)差和灰度圖不一樣工程里一般會提前計算訓(xùn)練集的通道均值后再歸一化??傊ㄓ霉こ涛募恼軐W(xué)是把不同任務(wù)里相同的那部分抽出來把差異化的那部分通過配置暴露出來。你訓(xùn)練的是手寫數(shù)字但復(fù)用的是工程框架。在這個基礎(chǔ)上每當(dāng)你要接一個新任務(wù)要改的只有數(shù)據(jù)和模型定義訓(xùn)練、評估、保存、推理那套鏈路幾乎不用動。這才是我理解的“完整工程文件”的意義。本文還有配套的精品資源點擊獲取