數(shù)字識(shí)別實(shí)戰(zhàn):從數(shù)據(jù)加載到CNN訓(xùn)練與模型部署)
簡(jiǎn)介MNIST手寫(xiě)數(shù)字識(shí)別是深度學(xué)習(xí)入門的經(jīng)典任務(wù)這份資源面向AI初學(xué)者與圖像識(shí)別從業(yè)者以TensorFlow構(gòu)建卷積神經(jīng)網(wǎng)絡(luò)完成09數(shù)字分類并附帶訓(xùn)練好的模型權(quán)重省去重復(fù)訓(xùn)練成本。壓縮包共6個(gè)文件、約2.19MB包含兩個(gè)Python腳本、可加載的h5預(yù)訓(xùn)練權(quán)重、訓(xùn)練過(guò)程可視化圖、預(yù)測(cè)效果示例圖及說(shuō)明文檔腳本分工明確便于直接運(yùn)行或二次修改。資源覆蓋數(shù)據(jù)預(yù)處理、模型搭建、訓(xùn)練評(píng)估和保存加載等關(guān)鍵環(huán)節(jié)通過(guò)加載權(quán)重即可快速測(cè)試識(shí)別效果并結(jié)合損失/準(zhǔn)確率曲線觀察收斂過(guò)程目前已有4621人學(xué)習(xí)。對(duì)想快速上手CNN的讀者來(lái)說(shuō)既能從零查看訓(xùn)練代碼也能利用預(yù)訓(xùn)練權(quán)重跳過(guò)訓(xùn)練階段說(shuō)明文檔中的提示還可幫助排查環(huán)境配置等問(wèn)題是理解特征提取與參數(shù)調(diào)優(yōu)的良好參考。1. MNIST 手寫(xiě)數(shù)字識(shí)別深度學(xué)習(xí)入門繞不開(kāi)的“第一課”很多剛開(kāi)始學(xué)深度學(xué)習(xí)的人第一個(gè)真正跑通并出結(jié)果的項(xiàng)目就是 MNIST 手寫(xiě)數(shù)字識(shí)別。這個(gè)數(shù)據(jù)集由 60000 張訓(xùn)練圖和 10000 張測(cè)試圖組成每張都是 28×28 的灰度手寫(xiě)數(shù)字配合卷積神經(jīng)網(wǎng)絡(luò)CNN你能在很短的時(shí)間內(nèi)親眼看到“圖像從像素變成高維向量再?gòu)南蛄孔兂煞诸惤Y(jié)果”的完整鏈路。它能解決的不只是一道課后題小尺寸、單通道、類別固定意味著你不需要先去折騰大規(guī)模分布式訓(xùn)練也能把注意力放在數(shù)據(jù)加載、反向傳播、準(zhǔn)確率調(diào)優(yōu)這些關(guān)鍵環(huán)節(jié)上。無(wú)論你是準(zhǔn)備系統(tǒng)學(xué) PyTorch還是想驗(yàn)證自己深度學(xué)習(xí)環(huán)境的可用性這套組合都是最直接的驗(yàn)證手段。這篇文章我按自己動(dòng)手的路徑來(lái)寫(xiě)先講數(shù)據(jù)長(zhǎng)什么樣再給一套能直接跑的訓(xùn)練代碼然后把最容易翻車的幾個(gè)坑一次說(shuō)清最后聊聊訓(xùn)練好的模型文件怎么保存、加載和復(fù)用。2. 先看清 MNIST 數(shù)據(jù)文件格式、加載方式和預(yù)處理2.1 一張手寫(xiě)數(shù)字圖像在硬盤上怎么存的MNIST 不是一堆 PNG 或 JPG 圖片而是采用 IDX 二進(jìn)制格式。每張圖固定是 28×28×8bit也就是 784 個(gè)灰度像素取值范圍 0255。你從官方渠道拿到手的是四個(gè)二進(jìn)制文件訓(xùn)練圖像、訓(xùn)練標(biāo)簽、測(cè)試圖像、測(cè)試標(biāo)簽。前 16 個(gè)字節(jié)是文件頭包含魔數(shù)和各維度信息第 16 字節(jié)之后每 784 字節(jié)就是一張完整的圖像按行優(yōu)先展開(kāi)。理解這個(gè)格式的好處在于當(dāng)你不想依賴 torchvision 時(shí)可以自己寫(xiě)幾行 Python 把原始數(shù)據(jù)讀出來(lái)。很多入門書(shū)和吳恩達(dá)課程的配套練習(xí)也都是圍繞這個(gè)二進(jìn)制格式展開(kāi)的。這里有段我常用來(lái)做“拆包驗(yàn)證”的代碼能順便確認(rèn)你下載的文件沒(méi)有損壞import numpy as np def load_mnist_images(path): # 讀取 MNIST 原始 IDX 文件返回形狀為 [N, 28, 28] 的 uint8 數(shù)組 with open(path, rb) as f: data f.read() magic int.from_bytes(data[:4], big) # 魔數(shù)校驗(yàn)文件類型 n int.from_bytes(data[4:8], big) # 圖像數(shù)量 rows int.from_bytes(data[8:12], big) # 高度 cols int.from_bytes(data[12:16], big) # 寬度 imgs np.frombuffer(data[16:], dtypenp.uint8).reshape(n, rows, cols) return imgs讀出來(lái)的 imgs 是 [60000, 28, 28] 的矩陣后面你想畫(huà)圖、做可視化或喂給自定義網(wǎng)絡(luò)都很方便。這里有個(gè)容易忽略的細(xì)節(jié)IDX 文件里所有整數(shù)都是大端big-endian存儲(chǔ)不能用 intel 小端的默認(rèn)方式直接解析否則讀出來(lái)的維度會(huì)亂得離譜。解析頭部四個(gè) int32 是最容易踩的底層坑一旦魔數(shù)不對(duì)先懷疑字節(jié)序。2.2 用 torchvision 直接下載并預(yù)處理如果是正式訓(xùn)練我推薦直接用 torchvision.datasets.MNIST。它幫你封裝好了下載、拆包、標(biāo)簽映射這些瑣碎邏輯也能直接對(duì)每一張 PIL 圖像做變換。最關(guān)鍵的一步是歸一化先用 ToTensor 把 0255 的像素縮放到 01再用 Normalize 按通道做標(biāo)準(zhǔn)化。MNIST 是灰度圖所以均值只有一個(gè)標(biāo)準(zhǔn)差也只有一個(gè)。import torch from torchvision import datasets, transforms # 這兩個(gè)標(biāo)準(zhǔn)化參數(shù)來(lái)自 MNIST 全部訓(xùn)練像素的統(tǒng)計(jì)值不是隨口填的 transform transforms.Compose([ transforms.ToTensor(), # [0,255] 像素 - [0,1] 浮點(diǎn)張量形狀 [1,28,28] transforms.Normalize((0.1307,), (0.3081,)) # 減均值再除標(biāo)準(zhǔn)差 ]) train_set datasets.MNIST(root./data, trainTrue, transformtransform, downloadTrue) test_set datasets.MNIST(root./data, trainFalse, transformtransform, downloadTrue)ToTensor 這一步很容易被誤解它不只是改數(shù)據(jù)類型還會(huì)把 HWC 的 PIL 圖像轉(zhuǎn)成 CHW 的張量MNIST 的通道數(shù)為 1所以最終形狀是 [1, 28, 28]。Normalize 不是可選操作跳過(guò)它訓(xùn)練也能收斂但收斂速度和最終精度都會(huì)受影響因?yàn)闆](méi)做標(biāo)準(zhǔn)化的輸入會(huì)讓網(wǎng)絡(luò)第一層梯度分布不穩(wěn)定。這套 (0.1307, 0.3081) 是社區(qū)里通用的 MNIST 統(tǒng)計(jì)值直接拿來(lái)用沒(méi)有毛病。2.3 DataLoader 的參數(shù)怎么設(shè)數(shù)據(jù)準(zhǔn)備到這一步還差一個(gè) DataLoader。這里有幾個(gè)參數(shù)直接決定訓(xùn)練體驗(yàn)from torch.utils.data import DataLoader train_loader DataLoader( train_set, batch_size64, # 每批 64 張圖梯度更平穩(wěn) shuffleTrue, # 每個(gè) epoch 都打亂順序避免模型學(xué)到樣本順序 num_workers4, # 多進(jìn)程加載Windows 上偶爾需要設(shè)為 0 drop_lastFalse # 最后一批不夠 64 張也保留不影響 MNIST ) test_loader DataLoader( test_set, batch_size256, shuffleFalse, # 評(píng)估時(shí)不需要打亂 num_workers4 )shuffleTrue 必須寫(xiě)在訓(xùn)練集上。很多新手在這翻車不打亂樣本順序模型會(huì)先看到大量類別 0再看到大量類別 1訓(xùn)練前期 loss 會(huì)規(guī)律性震蕩而且測(cè)試準(zhǔn)確率會(huì)卡在一個(gè)偏低的水平。num_workers 在 Linux 下開(kāi) 4 或 8 問(wèn)題不大Windows 下如果報(bào) DataLoader worker 相關(guān)錯(cuò)誤先降到 0 排查。3. 用卷積神經(jīng)網(wǎng)絡(luò)訓(xùn)練手寫(xiě)數(shù)字識(shí)別模型網(wǎng)絡(luò)設(shè)計(jì)與完整訓(xùn)練代碼3.1 為什么選卷積神經(jīng)網(wǎng)絡(luò)而不是前饋全連接網(wǎng)絡(luò)前饋神經(jīng)網(wǎng)絡(luò)FNN處理圖像時(shí)通常要把 28×28 展平成 784 維向量這等于強(qiáng)行丟掉像素之間的二維空間關(guān)系。數(shù)字 7 的橫線和豎線在展開(kāi)成一維后可能相隔幾百個(gè)位置全連接層需要自己硬學(xué)出這種遠(yuǎn)距離相關(guān)性參數(shù)量大而且樣本效率低。卷積神經(jīng)網(wǎng)絡(luò)不一樣它用固定大小的卷積核在圖像上滑動(dòng)天然只關(guān)注局部窗口同時(shí)又通過(guò)堆疊層數(shù)逐步擴(kuò)大感受野。對(duì) MNIST 來(lái)說(shuō)第一層卷積往往能學(xué)到橫線、豎線、斜線這類基礎(chǔ)筆畫(huà)第二層組合出環(huán)、拐角、交叉點(diǎn)這些結(jié)構(gòu)越往后越接近“數(shù)字部件”的抽象表示。這正好解釋了為什么一個(gè)只有幾萬(wàn)參數(shù)的輕量 CNN也能在測(cè)試集上拿到 99% 級(jí)別的準(zhǔn)確率。反觀全連接網(wǎng)絡(luò)同樣參數(shù)量下通常要低一到兩個(gè)百分點(diǎn)而且訓(xùn)練更慢。3.2 一個(gè)能直接跑的輕量 CNN 結(jié)構(gòu)下面這個(gè)結(jié)構(gòu)是典型的“卷積-池化-卷積-池化-全連接”鏈路參數(shù)規(guī)模很小CPU 上幾分鐘就能訓(xùn)完GPU 上更快。后續(xù)你想換成 ResNet 或更深的網(wǎng)絡(luò)也是從這個(gè)骨架長(zhǎng)出來(lái)的。import torch.nn as nn class SimpleCNN(nn.Module): def __init__(self): super().__init__() # 輸入 1 個(gè)通道輸出 32 個(gè)特征圖3x3 卷積padding1 保持尺寸 28x28 self.conv1 nn.Conv2d(1, 32, kernel_size3, padding1) # 第二層加深到 64 個(gè)特征圖 self.conv2 nn.Conv2d(32, 64, kernel_size3, padding1) # 2x2 最大池化特征圖從 28 - 14 - 7 self.pool nn.MaxPool2d(kernel_size2, stride2) # 最后一層池化后是 64 個(gè) 7x7 特征圖展平后正好 3136 維 self.fc1 nn.Linear(64 * 7 * 7, 128) self.fc2 nn.Linear(128, 10) # 10 個(gè)數(shù)字類別 def forward(self, x): x torch.relu(self.conv1(x)) x self.pool(x) x torch.relu(self.conv2(x)) x self.pool(x) x x.view(x.size(0), -1) # 展平成 [batch, 3136] x torch.relu(self.fc1(x)) x self.fc2(x) # 不接 softmax損失函數(shù)內(nèi)部處理 return x這個(gè)網(wǎng)絡(luò)總計(jì)約 42 萬(wàn)參數(shù)對(duì) MNIST 這種簡(jiǎn)單任務(wù)來(lái)說(shuō)已經(jīng)超過(guò)了“夠用”的線。padding1 是為了讓卷積不縮小尺寸這樣下采樣完全交給池化層做維度變化可預(yù)測(cè)。如果你去掉 padding28×28 會(huì)在第一層直接變成 26×26后面全連接層的輸入維度就要重新算。最后的全連接層不接 softmax是因?yàn)?PyTorch 的 CrossEntropyLoss 內(nèi)部會(huì)先算 softmax 再算交叉熵你提前 softmax 反而會(huì)數(shù)值不穩(wěn)。3.3 訓(xùn)練循環(huán)與關(guān)鍵超參數(shù)訓(xùn)練循環(huán)的核心是四個(gè)步驟梯度清零、前向傳播、計(jì)算損失、反向傳播。每次迭代都按這個(gè)順序來(lái)一步都不能亂。下面是完整可用的訓(xùn)練代碼import torch import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms device torch.device(cuda if torch.cuda.is_available() else cpu) # 數(shù)據(jù)準(zhǔn)備 transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_set datasets.MNIST(root./data, trainTrue, transformtransform, downloadTrue) train_loader DataLoader(train_set, batch_size64, shuffleTrue, num_workers4) # 模型、損失函數(shù)、優(yōu)化器 model SimpleCNN().to(device) criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr1e-3) # 訓(xùn)練 5 個(gè) epoch 通常足夠 for epoch in range(5): model.train() # 進(jìn)入訓(xùn)練模式啟用 Dropout/BatchNorm 的訓(xùn)練行為 running_loss, correct, total 0.0, 0, 0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() # 清空上一批的梯度 outputs model(images) # 前向傳播像素 - logits 向量 loss criterion(outputs, labels) # 計(jì)算交叉熵?fù)p失 loss.backward() # 反向傳播計(jì)算梯度 optimizer.step() # 更新權(quán)重 running_loss loss.item() * images.size(0) correct (outputs.argmax(1) labels).sum().item() total images.size(0) print(fepoch {epoch1:2d} | loss {running_loss/total:.4f} | acc {correct/total:.4f})我一般用如下超參組合對(duì)新手最不容易翻車參數(shù)取值說(shuō)明batch_size64太小梯度抖動(dòng)太大收斂變慢學(xué)習(xí)率1e-3Adam 的默認(rèn)學(xué)習(xí)率多數(shù)場(chǎng)景直接能用epochs5MNIST 上 5 輪足夠逼近 99%優(yōu)化器Adam對(duì)學(xué)習(xí)率不那么敏感適合快速驗(yàn)證損失函數(shù)CrossEntropyLoss內(nèi)部完成 softmax 交叉熵這里的梯度清零是很多新手容易漏的。PyTorch 的梯度默認(rèn)是累積的你不在每個(gè) batch 前調(diào)用 zero_grad下一輪反向傳播會(huì)把新舊梯度加在一起導(dǎo)致 loss 曲線異常震蕩。如果你觀察到 loss 持續(xù)下降但突然跳高先排查是不是忘了清零。4. MNIST 上手項(xiàng)目的五個(gè)高頻踩坑記錄與排查清單4.1 現(xiàn)象torchvision 下載 MNIST 時(shí) 404 或長(zhǎng)期卡住原因torchvision 的 MNIST 下載地址指向國(guó)外服務(wù)器國(guó)內(nèi)網(wǎng)絡(luò)環(huán)境經(jīng)常超時(shí)或返回 404。這不是代碼問(wèn)題是網(wǎng)絡(luò)問(wèn)題。解決手動(dòng)從可訪問(wèn)的鏡像下載四個(gè)文件然后放進(jìn) root 目錄。MNIST 文件是固定的四個(gè)train-images、train-labels、test-images、test-labels。放好后把 download 設(shè)為 False就不會(huì)再觸發(fā)遠(yuǎn)程下載。train_set datasets.MNIST(root./data, trainTrue, downloadFalse) # 如果目錄里已有對(duì)應(yīng)文件downloadFalse 會(huì)直接加載本地?cái)?shù)據(jù)4.2 現(xiàn)象loss 下降但測(cè)試準(zhǔn)確率卡在 90% 附近原因最常見(jiàn)是 DataLoader 沒(méi)開(kāi) shuffle或者輸入圖片沒(méi)做歸一化。前一種情況會(huì)讓模型學(xué)到樣本順序的虛假規(guī)律后一種會(huì)讓網(wǎng)絡(luò)權(quán)重更新路徑不穩(wěn)定。解決先確認(rèn) train_loader 里 shuffleTrue再檢查 transform 里有沒(méi)有 ToTensor 和 Normalize。如果兩個(gè)都正常還不漲把學(xué)習(xí)率從 1e-3 調(diào)低到 3e-4 再試。90% 這個(gè)位置通常是“模型在學(xué)但沒(méi)學(xué)好”的信號(hào)而不是模型容量不夠。4.3 現(xiàn)象訓(xùn)練集準(zhǔn)確率 99% 以上測(cè)試集卻明顯落后原因過(guò)擬合。網(wǎng)絡(luò)把訓(xùn)練樣本的噪聲也記進(jìn)去了尤其當(dāng)全連接層參數(shù)太多時(shí)這個(gè)現(xiàn)象非常明顯。解決在全連接層之間加 Dropout推理時(shí)它自動(dòng)關(guān)閉不需要額外處理。加了 Dropout 后測(cè)試準(zhǔn)確率通常能回升而且訓(xùn)練準(zhǔn)確率稍微降一點(diǎn)是正常的說(shuō)明模型不再死記硬背。self.dropout nn.Dropout(0.5) # 50% 概率隨機(jī)丟棄神經(jīng)元 # forward 里在 fc1 和 fc2 之間加一行 x self.dropout(torch.relu(self.fc1(x)))4.4 現(xiàn)象同一個(gè)代碼跑兩遍結(jié)果不完全一樣原因深度學(xué)習(xí)訓(xùn)練本身的隨機(jī)性包括權(quán)重初始化、DataLoader 打亂順序、GPU 上的并行計(jì)算順序。這不是玄學(xué)是默認(rèn)行為。解決在訓(xùn)練腳本最開(kāi)頭固定隨機(jī)種子。def set_seed(seed42): import random import numpy as np random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.backends.cudnn.deterministic True set_seed(42)4.5 現(xiàn)象把自己手寫(xiě)的圖片送進(jìn)模型預(yù)測(cè)結(jié)果離譜原因幾乎都是預(yù)處理不一致。訓(xùn)練時(shí)用的是 28×28 灰度圖 ToTensor Normalize推理時(shí)直接拿一張 RGB 彩色圖或原始尺寸圖片喂進(jìn)去模型當(dāng)然不認(rèn)。解決推理前必須對(duì)輸入做完全相同的一串變換。用 PIL 讀圖后先轉(zhuǎn)灰度再縮放到 28×28最后套同一套 transform。曾經(jīng)有人因?yàn)橥?resize 而讓 224×224 的圖片直接進(jìn)模型準(zhǔn)確率跌得比隨機(jī)猜還慘。5. 把訓(xùn)練好的模型文件用起來(lái)保存、加載與推理驗(yàn)證5.1 只存 state_dict 還是存整個(gè)模型常見(jiàn)做法是只保存 state_dict也就是網(wǎng)絡(luò)的參數(shù)權(quán)重不保存模型結(jié)構(gòu)。這樣做文件體積小而且換機(jī)器加載時(shí)只要用相同結(jié)構(gòu)的類實(shí)例化就能恢復(fù)。完整保存整個(gè)模型雖然省事但 PyTorch 版本一換經(jīng)常出兼容性問(wèn)題所以我一般不用。torch.save(model.state_dict(), ./mnist_cnn.pth)加載時(shí)要注意兩點(diǎn)第一必須先用類創(chuàng)建出模型實(shí)例第二加載后要調(diào)用 eval() 切換到推理模式。很多人漏掉 eval()結(jié)果推理結(jié)果時(shí)好時(shí)壞因?yàn)?Dropout 在訓(xùn)練模式下還會(huì)隨機(jī)丟棄神經(jīng)元。model SimpleCNN() # 先構(gòu)造相同結(jié)構(gòu) model.load_state_dict(torch.load(./mnist_cnn.pth, map_locationcpu)) model.eval() # 切換到推理模式關(guān)閉 Dropout 的隨機(jī)行為5.2 在測(cè)試集上驗(yàn)證模型文件是否真的可用拿到別人給你的模型文件不要直接拿去部署先在測(cè)試集上驗(yàn)證一遍。這招能幫你識(shí)別文件損壞、結(jié)構(gòu)不匹配、預(yù)處理不一致三類問(wèn)題。下面這段代碼會(huì)輸出最終測(cè)試準(zhǔn)確率def evaluate(model, loader): model.eval() correct, total 0, 0 with torch.no_grad(): # 推理不需要計(jì)算梯度省內(nèi)存 for images, labels in loader: images, labels images.to(device), labels.to(device) preds model(images).argmax(1) correct (preds labels).sum().item() total labels.size(0) return correct / total test_acc evaluate(model, test_loader) print(ftest acc: {test_acc:.4f})with torch.no_grad() 不是可選項(xiàng)。如果推理時(shí)還保留梯度計(jì)算顯存占用會(huì)明顯上漲而且速度慢很多。如果驗(yàn)證準(zhǔn)確率只有 10% 左右基本是類別標(biāo)簽錯(cuò)位或者預(yù)處理沒(méi)對(duì)齊先別懷疑模型回頭檢查 transform。5.3 單張圖片推理與最后的實(shí)踐習(xí)慣單張推理和批量評(píng)估的差異只在 batch 維度。單張圖需要手動(dòng)加一個(gè)維度因?yàn)槟P推谕斎胧?[batch, channel, height, width]而單張圖只有 [channel, height, width]def predict_one(model, img_tensor, device): model.eval() with torch.no_grad(): logits model(img_tensor.unsqueeze(0).to(device)) return logits.argmax(1).item()img_tensor 必須先經(jīng)過(guò)和訓(xùn)練時(shí)一致的 transform否則前面第 4.5 節(jié)的坑會(huì)再次出現(xiàn)。我會(huì)加一句最終 logits 向量里的數(shù)值大小可以當(dāng)作一個(gè)粗糙的置信度參考如果最大值和次大值非常接近說(shuō)明模型在兩個(gè)類別之間猶豫這個(gè)樣本要么寫(xiě)得太潦草要么不在訓(xùn)練分布內(nèi)部署時(shí)要對(duì)這種情況設(shè)置人為的拒絕閾值。這個(gè)項(xiàng)目打到 99% 以上只能算“入門完成”但真正收尾的功夫在驗(yàn)證和保存的規(guī)范上。我自己的習(xí)慣是每次訓(xùn)練完把測(cè)試準(zhǔn)確率、用了幾個(gè) epoch、超參組合寫(xiě)成一行文本放在模型文件旁邊免得三天后再看模型時(shí)完全想不起來(lái)當(dāng)初怎么調(diào)出來(lái)的。吃了不少這樣的虧之后我再也不存“裸模型”了。朋友拿一個(gè)沒(méi)說(shuō)明的權(quán)重文件來(lái)找我跑第一步永遠(yuǎn)是先問(wèn)測(cè)試集準(zhǔn)確率和預(yù)處理方式因?yàn)檫@兩樣對(duì)不上模型文件就只是一堆無(wú)法使用的數(shù)字。希望這份手把手的流程能幫到你少走幾趟彎路。本文還有配套的精品資源點(diǎn)擊獲取