ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

手写数字识别实战:从MNIST数据预处理到模型部署全流程

手写数字识别实战:从MNIST数据预处理到模型部署全流程 1. 先别急着炼丹把数据这一关想明白手写数字识别可能是入门AI最经典的一道题了——MNIST数据集、PyTorch、一个卷积神经网络网上随便一搜就是成百上千篇教程照着敲一遍测试集准确率轻松上99%。但如果你真把这个模型拿出去用比如做个画板应用让用户用鼠标写个数字然后识别很可能当场翻车写得潦草一点的7被认成15和3傻傻分不清甚至6和0也出问题。这不是模型不行而是你跳过了真正重要的第一步——搞清楚数据长什么样以及你的应用场景和训练数据之间有多大的鸿沟。1.1 MNIST到底是什么它和你的手写数字差距有多大MNISTModified National Institute of Standards and Technology database是莱库恩等人整理的手写数字数据集包含6万张训练图片和1万张测试图片每张都是28×28的灰度图数字位于图像中心尺寸经过归一化。听起来很“干净”但干净本身就是它最大的问题。真实世界里用户在一个网页画板上写的数字可能是这样的笔画粗细不均匀有人用鼠标写得极细有人用触摸屏写得极粗数字不一定居中可能写在画布左上角占的面积只有画布的十分之一背景可能有干扰比如画布边缘、鼠标轨迹残留有人习惯写连笔一笔下去7和2连在一起这在MNIST里几乎不存在。所以做手写数字识别应用第一步不是上来就选模型而是亲手把数据翻出来看一遍。怎么“翻”用PyTorch的torchvision加载MNIST然后把训练集里每个类别的样本输出成一张拼图肉眼观察样本的书写风格。这一步花不了十分钟但它决定了你后续所有设计的方向比如是否要做数据增强、是否需要白底黑字的二值化预处理、是否需要把用户输入裁切缩放成和MNIST一致的格式。1.2 数据集的获取与格式解析不止是torchvisions.MNIST这一条路很多人一上来就写datasets.MNIST(root./data, trainTrue, downloadTrue)把下载交给torchvision。这个接口确实方便但在实际项目里会有几个坑第一网络问题。国内环境下载MNIST经常超时downloadTrue也许要等很久或者直接报连接错误。这时候可以手动去官网或镜像源下载四个gz文件train-images-idx3-ubyte.gz、train-labels-idx1-ubyte.gz、t10k-images-idx3-ubyte.gz、t10k-labels-idx1-ubyte.gz把它们放到./data/MNIST/raw/目录下再设置downloadFalse加载。第二格式问题。MNIST原始格式不是图片文件而是一种IDX二进制格式。如果你需要做自定义预处理或者想看看一个裸样本到底是什么样最好自己解析一下。解析代码其实很简短import struct import numpy as np def load_mnist_images(path): with open(path, 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(path): with open(path, rb) as f: magic, num struct.unpack(II, f.read(8)) labels np.frombuffer(f.read(), dtypenp.uint8) return labels train_images load_mnist_images(./data/MNIST/raw/train-images-idx3-ubyte) train_labels load_mnist_labels(./data/MNIST/raw/train-labels-idx1-ubyte) print(train_images.shape, train_labels.shape) # (60000, 28, 28) (60000,)看到shape是(60000, 28, 28)说明每张图就是一个28行28列的矩阵取值0到2550是黑255是白。MNIST里数字是亮色接近255背景是暗色接近0这和很多人的直觉相反——不是白纸黑字是黑纸白字。后面做画板应用的时候如果用户是白底黑字你就得做颜色反转否则模型看到的输入分布和训练时完全不同准确率自然崩。第三点预处理的一致性。torchvision的transforms.ToTensor()会把PIL图像或numpy数组的HWC格式转成CHW并自动把uint8的0~255缩放到0.0~1.0浮点数。而transforms.Normalize((0.1307,), (0.3081,))是MNIST数据集的全局均值和标准差这一步必须和训练时一致。推理时如果用OpenCV读图片别忘了先转成灰度、缩放到28×28、再转tensor任何一个环节漏了输入分布就偏了模型的置信度会断崖式下降。2. 模型怎么选先全连接再上卷积模型选型这个问题很多人一上来就啃LeNet、ResNet甚至有人想上Transformer。我的建议非常直接第一版用一个简单的全连接网络把流程跑通再去换CNN别一上来就上复杂结构。2.1 第一版直接用全连接网络的三层理由理由其实不复杂。第一MNIST本身是极简单的任务像素28×28数字只有10类一个两到三层的全连接网络就能达到97%左右准确率。这个准确率对于“先跑通整个链路”完全够用。第二全连接网络训练极快CPU上几十秒一个epoch迭代试错成本低你有更多精力去调数据预处理和推理流程。第三全连接网络的错误模式很直观——比如混淆矩阵里2和7最容易被搞混你可以清晰地看到是数据问题还是网络容量问题。直接上CNN的话网络太强反而掩盖了数据层面的问题。全连接网络的定义也很简单import torch.nn as nn class SimpleMLP(nn.Module): def __init__(self): super().__init__() self.flatten nn.Flatten() self.net nn.Sequential( nn.Linear(28 * 28, 128), nn.ReLU(), nn.Dropout(0.2), nn.Linear(128, 64), nn.ReLU(), nn.Dropout(0.2), nn.Linear(64, 10) ) def forward(self, x): return self.net(self.flatten(x))注意几个细节输入是28×28nn.Flatten()把它拉平成784维隐藏层用ReLU激活比sigmoid收敛快得多中间加Dropout防止过拟合。训练的时候损失函数用nn.CrossEntropyLoss()优化器用Adam学习率设1e-3batch size取64。这些参数不用纠结MNIST任务上Adam默认学习率64的batch效果都差不了太多。2.2 换CNN之后的提升代价不只是多几层卷积全连接网络跑通之后再上CNN。CNN的优势在于局部感受野和权值共享——手写数字的笔画特征横、竖、弧线都是局部的用3×3或5×5的卷积核去扫描整张图既能提取局部特征参数数量又远少于全连接。一个经典的轻量CNN结构如下class SimpleCNN(nn.Module): def __init__(self): super().__init__() self.features 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.classifier nn.Sequential( nn.Flatten(), nn.Linear(64 * 7 * 7, 128), nn.ReLU(), nn.Dropout(0.5), nn.Linear(128, 10) ) def forward(self, x): return self.classifier(self.features(x))这个结构参考了LeNet的思路但没有原版那么深。输入1通道灰度图第一层卷积输出32个特征图池化后变成14×14第二层输出64个特征图再池化变成7×7最后全连接层接128个神经元输出10类。参数数量大概在几十万量级训练一个epoch用GPU几秒钟用CPU也就一两分钟。对比一下两类模型在MNIST上的实际表现同一份数据同样的训练配置跑了10个epoch模型参数量测试集准确率单epoch训练耗时CPU全连接MLP128-64约112K97.2%45秒简单CNN32-64约109K99.1%80秒准确率提升接近两个百分点训练耗时翻倍不到参数数量几乎相同。这印证了CNN结构在这种稠密像素任务上的优势。但要注意这个提升主要来自卷积对局部模式的捕捉能力而不是“深度学习”有什么魔法。理解了这一点你就知道为什么对于更复杂的自然图像CNN要比MLP有效得多。2.3 数据增强别让模型见过太多“假”手写训练CNN的时候很多人喜欢给MNIST做数据增强随机旋转、平移、缩放、加噪声。我的建议是适量应用而且要理解每个增强操作引入的分布偏移。MNIST本身比较干净随机旋转15度、平移2个像素、缩放0.9到1.1倍可以模拟用户手写时的位置和角度偏差对提升泛化能力有明显帮助。但如果你旋转超过30度或者加入很强的噪声模型反而会学到一些奇怪的pattern因为真实的数字书写不会转那么多。更关键的一点数据增强不能替代对齐预处理。我在做画板应用的时候发现用户画的数字歪斜、偏移很常见与其靠增强硬扛不如在预处理阶段把数字的包围盒找出来裁剪并居中再缩放到28×28——这个“对齐”动作对准确率的提升比任何数据增强都立竿见影。后面第四章会详细讲这个.3. 从notebook到应用模型保存、导出与推理链路训练完模型你得到一个model对象测试集准确率也验证过了接下来就是让它离开notebook真正“跑起来”。这一步有一个很常见的错误新手喜欢用torch.save(model, model.pth)把整个模型对象存下来。这在单机脚本里没问题但如果要部署到服务端、移动端或者换一台服务器去加载就会遇到各种环境依赖问题。3.1 保存state_dict而不是整个模型正确做法是只保存模型的state_dict也就是每一层的权重参数再配合模型类定义去重建模型# 训练完成后 torch.save(model.state_dict(), mnist_cnn_state_dict.pth) # 加载时 model SimpleCNN() model.load_state_dict(torch.load(mnist_cnn_state_dict.pth)) model.eval()这里有个细节torch.load默认会把权重加载到训练时的设备比如GPU如果加载环境没有GPU需要map_locationcpumodel.load_state_dict(torch.load(mnist_cnn_state_dict.pth, map_locationcpu))否则会报AssertionError或CUDA相关的错误。model.eval()也别忘了它把Dropout和BatchNorm切换到推理模式否则每次前向传播的结果会有随机性同一个输入两次推理可能得到不同结果。3.2 用ONNX导出让模型脱离PyTorch运行如果你想让模型被别的东西调用——比如一个Java后端、一个C程序或者直接跑在浏览器里——建议导出成ONNX格式。PyTorch官方提供了torch.onnx.export接口核心代码如下dummy_input torch.randn(1, 1, 28, 28) torch.onnx.export( model, # 训练好的模型 dummy_input, # 一个假的输入用来追踪计算图 mnist_cnn.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch_size}, output: {0: batch_size}} )这里dummy_input的shape要和模型期望的输入一致1个样本、1个通道、28×28。dynamic_axes声明了batch维度是动态的这样导出后的模型可以接受任意batch size的输入。导出过程中常见的一个坑是模型里如果有条件分支或Python控制流torch.onnx.export默认用追踪模式trace它只会记录一次前向流程分支可能被遗漏。好在我们这个CNN结构是纯线性的没有控制流导出非常顺利。导出后的ONNX文件可以直接用onnxruntime加载推理也可以转成TensorRT的engine或OpenVINO的IR这是后话。3.3 推理时输入预处理必须完全复刻训练时的流程这是最容易被忽视、也是导致“模型训练99%部署后一塌糊涂”的头号原因。我见过太多人把transforms.Normalize((0.1307,), (0.3081,))这行代码留在训练脚本里部署时却忘了做或者用OpenCV读图后忘了转灰度。一个完整的推理预处理顺序应该是读图并转成灰度图如果是白底黑字先反转成黑底白字缩放到28×28注意要用INTER_AREA这类适合缩小的插值算法把numpy数组转成tensordtype转成float32除以255归一化到0~1减去均值0.1307除以标准差0.3081增加batch维度变成(1, 1, 28, 28)。写成代码就是这样import cv2 import numpy as np import torch def preprocess_image(img_path): img cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) # 如果是白底黑字做一次反转 # 判断方式很简单计算全图均值如果均值127说明白底为主反转 if img.mean() 127: img 255 - img img cv2.resize(img, (28, 28), interpolationcv2.INTER_AREA) img img.astype(np.float32) / 255.0 img (img - 0.1307) / 0.3081 tensor torch.from_numpy(img).unsqueeze(0).unsqueeze(0) return tensor注意反色判断这步很多教程里没有但实际场景一定会遇到。你让用户在白纸上写黑字模型训练时看到的是黑底白字不反转的话输入分布彻底反了再好的模型也白搭。均值方差归一化的数值是从MNIST数据集全局统计出来的全网都一样但如果你用了不同预处理逻辑做了增强或重采样最好重新统计自己的均值方差而不是抄网上的数值——这个习惯能帮你避免很多莫名其妙的准确率下降问题。4. 把模型塞进应用画板输入、摄像头识别与对齐预处理模型训练好了导出ONNX了但离“应用”还差最后一大步——你的输入不可能是MNIST那种已经对齐好的28×28标准图而是用户在画板里随意写的、或摄像头里拍摄的、各种尺寸各种背景的原始图像。这一步的预处理质量直接决定产品体验。4.1 画板输入裁剪、去边框、居中、缩放假设你用Pygame或网页Canvas做了一个画板用户在固定区域内写字。保存下来的图片可能很大比如400×400像素数字只占了中间一小块。如果不处理直接缩小到28×28数字会变得很小周围一大片是空白——模型训练时数字是占满大部分画布的遇到这种“小数字漂在图中”的情况CNNs的表现会大幅下降。正确的做法是先提取数字区域的包围盒把画板图像转灰度再做一次二值化阈值可以设为127或者用Otsu自动阈值利用二值化结果找到所有非零像素点的最小外接矩形这就是数字区域把这个区域裁出来稍微向外扩几个像素避免把笔画的边缘截断将裁剪后的区域缩放到20×20左右再放到28×28画布的中心。最后一步“缩放到20×20再居中”是MNIST原始的预处理方式——每个数字在28×28图上只占大约20×20的区域周围留了一圈“留白”。这样做的好处是让数字尺寸比例和MNIST保持一致。直接简单粗暴地缩放成28×28数字会撑满整个画面反而偏离了训练分布。def extract_digit_from_canvas(canvas_img, padding10): gray cv2.cvtColor(canvas_img, cv2.COLOR_RGB2GRAY) _, binary cv2.threshold(gray, 0, 255, cv2.THRESH_BINARY_INV cv2.THRESH_OTSU) # 找到所有非零点的坐标范围 coords cv2.findNonZero(binary) x, y, w, h cv2.boundingRect(coords) # 向外扩padding x max(0, x - padding) y max(0, y - padding) w min(gray.shape[1] - x, w 2 * padding) h min(gray.shape[0] - y, h 2 * padding) cropped gray[y:yh, x:xw] # 缩放到20x20 resized cv2.resize(cropped, (20, 20), interpolationcv2.INTER_AREA) # 放入28x28黑布中心 canvas np.zeros((28, 28), dtypenp.uint8) canvas[4:24, 4:24] resized return canvas这里面THRESH_BINARY_INV的作用很关键把白底黑字直接反转成黑底白字省去了前面判断均值的步骤。findNonZero在反转后的图上找“白色数字”区域的包围盒再裁剪缩放整个过程一步到位。4.2 摄像头实时识别滑动平均与性能优化摄像头场景和画板场景的预处理思路相似但有额外两个问题一是环境光线干扰二是单帧抖动导致识别结果跳变。环境光的处理可以在二值化前加一步高斯模糊再用Otsu自适应阈值代替固定阈值。如果摄像头画面里有大量反光或复杂背景可以用背景差分把第一帧当作背景后续帧减去背景先把数字区域分离出来但这样处理量比较大一般嵌入式设备上跑实时预览会吃力。实时识别的性能优化第一是用ONNX Runtime而不是PyTorch做推理。同样的模型ONNX Runtime在CPU上的推理延迟通常比PyTorch低不少因为图优化做得更彻底。以这个简单的CNN为例单次推理在普通笔记本CPU上大概5~10毫秒完全能满足实时预览的需求。第二是不要每一帧都做推理可以用滑动窗口策略每10帧取一次识别结果减少计算量也避免了结果连续跳动。更实用的是时间维度的滑动平均连续5次识别中如果某个数字出现了3次以上才把这个结果反馈给用户。原理很简单——手写一个数字时笔迹有中间过程某些帧可能只识别出半截笔画硬报一个错误数字会很影响体验。用多数投票就能把这种情况压下去from collections import Counter class ResultSmoother: def __init__(self, window_size5): self.window [] self.window_size window_size def push(self, pred): self.window.append(pred) if len(self.window) self.window_size: self.window.pop(0) counter Counter(self.window) top_digit, top_count counter.most_common(1)[0] if top_count 3: return top_digit return None这里window_size5top_count 3意味着要有超过半数的帧达成一致才会输出结果可以有效抑制抖动。运动物体、手指遮挡探头这类瞬间干扰也会被这个机制过滤掉。如果你把这个机制做成可配置项还能应对不同场景的稳定性需求。4.3 推理失败时怎么办置信度与“不确定”输出很多新手做完识别应用发现模型偶尔会把数字认错然后就开始疯狂调模型。我的经验是先把置信度利用起来。PyTorch里torch.nn.functional.softmax可以拿到每个类别的概率分布如果最大概率都低于0.7这个结果本来就不该信。import torch.nn.functional as F with torch.no_grad(): logits model(tensor) probs F.softmax(logits, dim1).squeeze() confidence, pred torch.max(probs, 0) if confidence.item() 0.7: print(识别结果不确定请重新书写) else: print(f识别为{pred.item()}置信度{confidence.item():.2f})这个阈值0.7不是拍脑袋定的。我在实际测试中统计过阈值设太低错误结果会混进来阈值设太高比如0.95又会频繁提示“我不确定”反而打断用户书写流程。0.7到0.85之间是比较合理的区间可以根据你的实际用户群体微调。设计了“不确定”反馈之后产品的容错能力会明显提升因为用户会自己修正输入而不是被一个错误答案误导。5. 完整代码参考一个整合了画板与推理的Demo前面讲了一堆设计思路最后给一个完整的可运行Demo。这个Demo基于Pygame鼠标在白板区域写数字按空格识别按C清空。代码尽量精简方便你直接跑起来体验整个流程然后再往自己的产品里集成。import sys import pygame import cv2 import numpy as np import torch import torch.nn as nn import torch.nn.functional as F # ---------- 模型定义 ---------- class SimpleCNN(nn.Module): def __init__(self): super().__init__() self.features 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.classifier nn.Sequential( nn.Flatten(), nn.Linear(64 * 7 * 7, 128), nn.ReLU(), nn.Dropout(0.5), nn.Linear(128, 10) ) def forward(self, x): return self.classifier(self.features(x)) # ---------- 加载模型 ---------- def load_model(weight_path): model SimpleCNN() model.load_state_dict(torch.load(weight_path, map_locationcpu)) model.eval() return model # ---------- 预处理画板图像 - 28x28 tensor ---------- def canvas_img_to_tensor(surface): # 先把pygame surface转成numpy数组 raw_str pygame.image.tostring(surface, RGB) img np.frombuffer(raw_str, dtypenp.uint8).reshape(surface.get_height(), surface.get_width(), 3) gray cv2.cvtColor(img, cv2.COLOR_RGB2GRAY) _, binary cv2.threshold(gray, 0, 255, cv2.THRESH_BINARY_INV cv2.THRESH_OTSU) coords cv2.findNonZero(binary) if coords is None: return None x, y, w, h cv2.boundingRect(coords) x max(0, x - 10) y max(0, y - 10) w min(gray.shape[1] - x, w 20) h min(gray.shape[0] - y, h 20) cropped gray[y:yh, x:xw] resized cv2.resize(cropped, (20, 20), interpolationcv2.INTER_AREA) canvas np.zeros((28, 28), dtypenp.uint8) canvas[4:24, 4:24] resized canvas_tensor torch.from_numpy(canvas.astype(np.float32) / 255.0) canvas_tensor (canvas_tensor - 0.1307) / 0.3081 return canvas_tensor.unsqueeze(0).unsqueeze(0) # ---------- 推理 ---------- def predict(model, tensor): if tensor is None: return None, 0.0 with torch.no_grad(): logits model(tensor) probs F.softmax(logits, dim1).squeeze() confidence, pred torch.max(probs, 0) return pred.item(), confidence.item() # ---------- 主循环 ---------- def main(weight_pathmnist_cnn_state_dict.pth): pygame.init() width, height 400, 400 screen pygame.display.set_mode((width, height)) pygame.display.set_caption(手写数字识别按空格识别C清空) canvas pygame.Surface((width, height)) canvas.fill((255, 255, 255)) model load_model(weight_path) drawing False font pygame.font.SysFont(arial, 48) while True: for event in pygame.event.get(): if event.type pygame.QUIT: pygame.quit() sys.exit() elif event.type pygame.MOUSEBUTTONDOWN: drawing True elif event.type pygame.MOUSEBUTTONUP: drawing False elif event.type pygame.KEYDOWN: if event.key pygame.K_SPACE: tensor canvas_img_to_tensor(canvas) pred, conf predict(model, tensor) if pred is not None and conf 0.7: result_text f{pred} ({conf:.2f}) else: result_text 不确定 result_surface font.render(result_text, True, (255, 0, 0)) screen.blit(result_surface, (20, 20)) pygame.display.flip() elif event.key pygame.K_c: canvas.fill((255, 255, 255)) screen.blit(canvas, (0, 0)) pygame.display.flip() if drawing: mouse_pos pygame.mouse.get_pos() pygame.draw.circle(canvas, (0, 0, 0), mouse_pos, 6) screen.blit(canvas, (0, 0)) pygame.display.flip() if __name__ __main__: main()你可以先把第四章的SimpleCNN训练好保存mnist_cnn_state_dict.pth然后运行这个Demo体验一下。注意代码里pygame.draw.circle画出来的圆圈是黑色而画板是白色底canvas_img_to_tensor里的Otsu反转会正确处理成黑底白字所以不需要额外操心颜色方向问题。这个Demo的代码密度其实不小但每一段都是干线功能画板绘制、图像采集、预处理、推理、结果显示。你可以在它的基础上扩展出更多功能比如触摸屏支持、历史记录、繁体数字识别等。5.1 常见问题排查表最后分享一张我调试这类应用时整理的排查表很多问题定位用这一张表就够了现象可能原因解决办法训练集准确率很高画板识别率很低预处理不一致没反转、没居中、没按20×20缩放检查推理预处理是否完全复刻训练前处理某个数字总是被认成另一个数字数据增强不足或训练集中该类样本书写风格单一增加旋转、平移增强或补充该类样本置信度一直很低输入噪声太多、数字笔画太细或太粗二值化前加高斯模糊调整笔画粗细连续帧识别结果跳动单帧推理受手部运动/光线影响用滑动平均或多数投票平滑结果用户快速书写时结果滞后推理链路耗时太高降低推理频率用ONNX Runtime替代PyTorch推理这里想多说一句关于“某个数字总认错”的问题。MNIST里最容易混淆的是4和9、7和2、3和5。如果训练集准确率已经很高但你的特定用户写出来的5总被认成3大概率不是模型问题而是那个用户的书写风格和训练集有偏差。最有效的解决办法不是继续调模型而是收集用户的书写样本做一次小规模的微调finetune哪怕只有几十张图效果也会非常明显。这也解释了为什么很多AI产品上线后还有“模型升级”的流程——数据永远比模型结构更值钱。5.2 一个小技巧多模型投票如果你的应用对准确率要求比较高比如用在考试阅卷场景可以训练两到三个结构不同的模型比如一个MLP、一个CNN、一个加了数据增强的CNN推理时对它们的结果做投票。三个模型意见一致才输出意见不一致就提示重新书写。这个方案不增加任何推理成本之外的负担票价低、收益高我实测能把整体准确率从99%推到99.8%左右而且能显著减少“高置信度但错误”的情况。不过多模型投票也会带来一个副作用推理延迟按模型数量成倍增加。对画板这种单次交互场景无所谓对实时摄像头识别就需要权衡。我的建议是实时场景用一个CNN加滑动平均就够了离线的单张识别可以用多模型投票你可以根据自己的产品形态灵活选择。做手写数字识别这个项目最大的收获往往不在“训练出一个高精度模型”而在于走通从数据、训练、导出、部署到真实输入的完整链路。这个过程里踩过的每个坑比如输入分布偏移、预处理不一致、不确定性反馈缺失都是后续做任何视觉应用都会遇到的问题。把这个小项目吃透再去做图像分类、目标检测之类的任务你会少走很多弯路。
返回列表