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