ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

手写数字识别系统实战:从MNIST训练到Flask部署的完整流程

手写数字识别系统实战:从MNIST训练到Flask部署的完整流程 简介基于Python的手写数字识别系统完整项目面向计算机相关专业准备毕业设计或课程作业的学生也适合想动手实践神经网络的新手。项目评审分98分源码均经过本地调试可正常运行难度适中能满足学习与使用需求。压缩包共28个文件约14.18MB主要包括Python源码、训练参数、手写数字数据、说明文档和结果图表。其中脚本覆盖数据加载、卷积池化、反向传播、激活函数、训练与参数保存等完整流程参数文件保存了十次训练的准确率记录可清晰看到准确率从66.28%逐步提升至96.98%的迭代过程数据文件提供标准手写数字样本图表展示识别效果与训练走势。目前已有70人浏览学习。这套项目既能作为毕业设计代码框架也能帮助理解BP与卷积神经网络在图像识别中的实现差异配合参数文件对比分析便于读者改进模型或二次开发。1. 手写数字识别毕设为什么这个题目最容易做崩在“系统”而不是模型“基于Python实现的手写数字识别系统”几乎是每年毕设里出现频率最高的题目之一。原因很直接MNIST 是入门数据集网上现成代码一大把随便跑个 CNN 就能到 99% 准确率。但我要先说一个反直觉的结论这个题目真正让人翻车的从来不是模型精度而是大多数人只实现了“手写数字识别”没实现“系统”。数据怎么组织、模型怎么保存和加载、画板输入为什么识别不出来、换一台电脑为什么跑不起来——这些才是答辩时被追问最多的点。我会按自己平时做毕设辅导的习惯把从数据准备、训练、部署到答辩避坑的整条链路拆开讲所有细节都是可复现的。2. 数据与模型选型完整数据不是只有 MNIST识别系统的地基怎么打2.1 数据集构成MNIST、额外样本与标签文件缺一不可先明确“完整数据”在这个题目里指什么。MNIST 本身包含 60000 张训练图片和 10000 张测试图片每张都是 28x28 的灰度图0-9 十个类别基本均衡。torchvision 一行代码就能下载看起来省事但很多人在毕设提交时才发现自己根本没有单独的数据目录——数据散落在缓存目录里标签也没有可视化验证过。这会给评审一个非常不好的印象你说这是“系统”但数据资产是隐形的。我一般建议项目里至少要有三块数据资产。第一MNIST 原始数据放在 data/MNIST 下保留官方划分第二从训练集里切出的验证集以及一份 labels.csv 记录图片文件名和类别对应关系第三额外的手写样本目录可以是你自己用画板写的一组数字、也可以是扫描的纸质数字照片这部分是给后期演示和扩展用的也让“深度学习模型遇到分布外数据会怎样”这个问题在答辩时有东西可聊。目录结构上我习惯按下面这种方式组织代码、数据、权重三块彻底分开mnist_demo/ ├── data/ │ ├── MNIST/ # torchvision 自动下载的官方数据 │ ├── split/ # 统一处理后的训练/验证/测试文件 │ └── labels.csv # 文件名与标签的映射表 ├── checkpoints/ # 训练出来的模型权重 ├── src/ │ ├── model.py # CNN 网络定义 │ ├── train.py # 训练与评估脚本 │ ├── predict.py # 单张图片预测 │ └── app.py # Flask 识别接口 ├── requirements.txt └── README.md这个结构看起来普通但对毕设很重要评审打开目录能一眼找到入口答辩时你说“数据在这里、模型在这里、入口在这里”会显得思路清楚。labels.csv 不需要做得复杂一张 CSV 只有三列 filename、label、split对应每一张图属于训练集还是验证集。别小看这份文件答辩时展示它比口头说“数据有 70000 张”有说服力得多后面做混淆矩阵和错误样本分析时也离不开它。还要提醒一个高频问题MNIST 下载依赖外网资源第一次跑代码卡在 Downloading 很久然后失败的情况非常多。常见做法是手动把四个 .gz 文件准备好放到 data/MNIST/raw 目录下再运行训练脚本torchvision 检测到文件存在就不会重新下载。这四个文件分别是训练图片、训练标签、测试图片、测试标签文件名带 idx3 和 idx1 字样识别度很高。这个坑几乎是一切的起点我选择在这里先说清楚。2.2 模型选型逻辑为什么 Softmax 基线不可省CNN 才是主模型模型选型上我见过两种极端一种人直接抄一个大号的 ResNet训练半天精度不升反降另一种人从逻辑回归开始最后也拿逻辑回归交差精度只有 92% 左右。这两种都不理想。做毕设的正确打开方式是“两段式”先用一个极简单的 Softmax 回归把数据加载、训练、评估的整条 pipeline 跑通确保代码链路没有暗病再把模型换成 CNN 追求精度。这个做法省下的调试时间非常多——很多学生在 Softmax 阶段就会暴露数据归一化写错、标签错位、DataLoader 参数不对这类问题如果在 CNN 上才第一次接触这些问题排错会难得多。下面对比一下这个题目里常见的三种模型选择精度和答辩价值都不同模型参数量级MNIST 测试集预期精度答辩价值Softmax 回归约 8 千约 92%讲清交叉熵与损失函数单隐藏层 MLP约 10 万约 97%讲清过拟合与 Dropout两层卷积 CNN约 8 万99% 以上讲清卷积、池化、全连接链路选 CNN 而不是更大的网络是因为 MNIST 的图片只有 28x28 灰度信息量有限ResNet 这类深度网络在这个数据上提升很小反而更容易过拟合、训练更慢。两层卷积的 CNN 参数量和单层 MLP 接近容量却更强是性价比最高的选择。框架上我优先推荐 PyTorch原因不是它比 TensorFlow 快而是它的张量形状检查和调试体验更接近 Python 直觉答辩时老师问“你这行代码在做什么”你可以直接打印形状讲出来不会把解释环节变成黑匣子。2.3 数据预处理与增强让模型在答辩现场不翻车的三个参数数据预处理看起来是代码里最不起眼的几行但它决定训练能不能收敛。第一个参数是归一化。MNIST 的标准做法是把像素值从 [0, 255] 缩放到 [0, 1]ToTensor 自动完成再用全局均值 0.1307 和标准差 0.3081 做标准化。这两个数字是 MNIST 训练集的官方统计量直接照用不要自己重新算否则和模型期望输入的分布对不上推理阶段会莫名掉精度。第二个参数是变换顺序。torchvision 的 transforms.Compose 里ToTensor 必须在 Normalize 之前因为 ToTensor 负责把 PIL 图像或 numpy 数组转成张量并缩放到 0-1Normalize 直接在张量上做减法除法。顺序写反会直接报错或者得到错误分布这是新手最容易犯的错。第三个参数是数据增强的强度。很多人一上来就加随机旋转 30 度、随机裁剪结果训练集精度上去了验证集反而掉——因为 MNIST 测试集本身就是居中、无旋转的标准手写体过度增强把训练分布拉离了测试分布。我实际项目中只保留三个轻量增强随机旋转 ±10 度、随机平移 ±2 像素、以 0.2 的概率加入轻微椒盐噪声。这三个参数的意义是模拟真实手写的轻微偏移但不会破坏数字的结构。想验证增强是否过度有个很直观的办法每轮训练后打印验证集精度如果训练集一路涨、验证集在某个 epoch 后开始回落说明增强或模型容量过了优先降低增强强度而不是加大正则化。3. 用 Python 实现手写数字识别PyTorch 训练代码与参数说明3.1 环境搭建Python 版本、依赖安装与一个常见的 numpy 坑先说环境。Python 版本我建议 3.8 到 3.10 之间不要在项目开始时直接上 Python 3.12 加最新版 torch——不是不能用而是网上绝大多数教程和答疑帖都基于 3.8/3.10遇到问题你能搜到答案。用 conda 建一个独立环境是最省心的避免把系统 Python 弄乱也方便答辩前复现环境conda create -n mnist python3.10 conda activate mnist pip install torch torchvision matplotlib numpy pandas scikit-learn flask这里有一个值得单独提醒的搭配问题如果你先装了最新版 numpy比如 2.x再装 torch某些 torch 版本导入时会报类似_ARRAY_API not found的错误原因是 torch 官方二进制依赖的 numpy 版本和系统里的不一致。解决顺序很简单先装 torch再装 numpy让 pip 自动把 numpy 降到兼容版本如果已经踩了就按 pip 提示执行pip install numpy2。提示如果不是必须用 GPU就装 CPU 版本。MNIST 的数据量和两层 CNN 的规模用 CPU 训练完全够快一轮二十秒左右没必要在答辩前跟 CUDA 驱动较劲。Windows 用户装 torch 时如果pip install torch下载很慢常见做法是去 PyTorch 官网按自己的 CUDA 版本复制对应的安装命令或者直接用 CPU 版本。装完之后用一行代码验证环境python -c import torch, torchvision; print(torch.__version__, torchvision.__version__)能打印出版本号环境这关就算过了。显卡用户还可以加一句print(torch.cuda.is_available())但请注意这个项目用 CPU 训练就够了我不建议在毕设里引入 GPU 带来的设备依赖问题后面避坑章节会专门讲。3.2 加载 MNIST 数据代码与数据流说明数据加载是整个项目里最不该出错的部分因为它的错误隐藏得很深——代码能跑但跑出来的精度不对。下面这段是我项目里的标准加载写法import torch from torch.utils.data import DataLoader, random_split from torchvision import datasets, transforms # 固定随机种子保证每次切分出来的验证集一致方便复现 torch.manual_seed(42) # 归一化的均值和标准差是 MNIST 官方统计量不要自己算 transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) # 下载并加载训练集二次运行时 downloadTrue 不会重复下载 train_set datasets.MNIST(root./data, trainTrue, downloadTrue, transformtransform) # 从 60000 张训练图里切 5000 张作为验证集剩下的 55000 张用于训练 train_set, val_set random_split(train_set, [55000, 5000])这段代码有几个参数值得解释。root./data是数据目录和第 2 章的目录结构对应downloadTrue只在第一次生效文件存在后会自动跳过。random_split的第二个参数是每个子集的大小列表加起来必须等于原数据集长度否则会报错。为什么一定要切验证集因为如果你用全部 60000 张去训练就没有一个机制在训练过程中判断模型是否过拟合、该在哪个 epoch 保存权重。很多“测试集 99% 但一上手就废”的模型问题就出在这里。DataLoader 的参数同样有讲究train_loader DataLoader(train_set, batch_size64, shuffleTrue, num_workers2) val_loader DataLoader(val_set, batch_size256, shuffleFalse, num_workers2) test_loader DataLoader( datasets.MNIST(root./data, trainFalse, downloadTrue, transformtransform), batch_size256, shuffleFalse, num_workers2 )batch_size64是 MNIST 上很稳的默认值太大如 512 会让每次参数更新方向过于平均收敛慢且容易停在平坦区太小如 16 则梯度噪声大训练震荡。shuffleTrue只用于训练集验证集和测试集必须保持顺序否则当你用混淆矩阵分析错误时顺序一变你会以为模型输出和标签对不上。num_workers2在 Linux 上能加速数据读取但在 Windows 上如果报多线程错误直接改成 0否则会看到莫名其妙的死锁或重复加载。Windows 用户还要注意只要 num_workers 大于 0训练脚本就应该用if __name__ __main__包住避免多进程递归启动。3.3 构建 CNN 并训练网络结构、损失函数与 20 轮训练配置模型结构我用的是一个经典的两层卷积 CNN可以把它理解为 LeNet-5 的现代化简化版但用了 ReLU 激活和 Dropoutimport torch.nn as nn class DigitCNN(nn.Module): def __init__(self): super().__init__() # 输入 1x28x28padding2 让卷积后尺寸保持 28x28 self.conv1 nn.Conv2d(1, 32, kernel_size3, padding2) self.conv2 nn.Conv2d(32, 64, kernel_size3, padding2) # 两次 2x2 池化后 28 - 14 - 7所以全连接输入是 64*7*7 self.fc1 nn.Linear(64 * 7 * 7, 128) self.fc2 nn.Linear(128, 10) self.dropout nn.Dropout(0.25) def forward(self, x): x torch.relu(self.conv1(x)) x torch.max_pool2d(x, 2) x torch.relu(self.conv2(x)) x torch.max_pool2d(x, 2) x x.view(x.size(0), -1) x torch.relu(self.fc1(x)) x self.dropout(x) return self.fc2(x)这里有一个参数容易被忽略第一个卷积层的padding2而不是常用的padding1。3x3 卷积在 padding1 时同样能把 28x28 保持为 28x28但我在最大池化后想保持 14 和 7 的整除关系所以统一用 padding2 简化计算。你完全可以用 padding1效果几乎没有差别关键是结构里的尺寸注释要更新别让自己在调试时算错维度。self.fc2(x)没有接 ReLU 也没有接 Softmax原因是我在训练时用nn.CrossEntropyLoss()做损失函数它内部已经包含了 Softmax 和 log 操作输出层直接给 logits 就可以了。这是 PyTorch 里非常基础但很多新手会犯的误用有人在 fc2 后面手动加一个 softmax再喂给交叉熵损失结果得到的是双重 Softmax精度莫名其妙掉一截还很难排查。训练循环的写法我固定成下面这样简单而且可以复用到其他分类项目criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr1e-3) for epoch in range(20): model.train() total_loss 0 for images, labels in train_loader: optimizer.zero_grad() outputs model(images) # logits, 形状 (64, 10) loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() if (epoch 1) % 5 0: model.eval() correct 0 with torch.no_grad(): for images, labels in val_loader: preds model(images).argmax(dim1) correct (preds labels).sum().item() acc correct / len(val_set) print(fepoch {epoch1}, loss {total_loss / len(train_loader):.4f}, val_acc {acc:.4f})参数说明lr1e-3配合 Adam 是 MNIST 上最不容易翻车的组合如果训练 5 轮内 loss 完全不动优先检查数据而不是调学习率epochs20对两层 CNN 来说已经能到 99%再多轮的收益很小长尾收益不如去做数据增强。model.eval()和with torch.no_grad()在验证时缺一不可前者关掉 Dropout 和 BatchNorm 的随机性后者关掉梯度计算少了任何一个验证精度都会浮动甚至偏高。训练过程中如果想看 loss 曲线用 matplotlib 把每轮的 mean loss 画出来就行别在每批次都画横轴太密集看不清趋势还白耗时间。3.4 模型评估与保存准确率、混淆矩阵和 checkpoint 命名习惯训练完成后用测试集做一次完整评估并保存一个“带元信息”的 checkpoint这一步很多人做得太粗糙from sklearn.metrics import confusion_matrix model.eval() all_preds, all_labels [], [] with torch.no_grad(): for images, labels in test_loader: preds model(images).argmax(dim1) all_preds.extend(preds.tolist()) all_labels.extend(labels.tolist()) acc sum(p t for p, t in zip(all_preds, all_labels)) / len(all_labels) print(ftest_acc: {acc:.4f}) cm confusion_matrix(all_labels, all_preds) print(cm) torch.save({ model_state: model.state_dict(), class_names: [str(i) for i in range(10)], transform: {mean: [0.1307], std: [0.3081]}, test_acc: acc, }, checkpoints/mnist_cnn.pt)为什么我强调“带元信息”而不是只存model.state_dict()因为推理端加载模型时你需要知道用什么预处理、什么类别名、这个模型的精度是多少。把这些信息一起封进 checkpoint会让predict.py和app.py的代码干净很多也不依赖网络定义文件是不是还在同一个路径下。这个习惯在答辩现场能救你一次后面避坑章节会具体说。checkpoint 的文件名我习惯带上精度例如mnist_cnn_99.12.pt好处是多个实验版本并存时你不会拿错权重。混淆矩阵打印出来后重点看对角线以外哪些格子数值大那是模型系统性犯错的地方比如 4 和 9、7 和 9。这一张图就是你答辩时分析模型短板的素材比单纯报一个准确率高一个层次。4. 把手写识别系统跑成完整应用Web 接口与可视化界面4.1 用 Flask 封装识别接口POST 图片返回数字与置信度训练出模型只完成了一半“系统”的完整度主要体现在能不能被别人用起来。我常用 Flask 包一个 HTTP 接口原因很简单它是标准库级轻量的方案答辩演示时只需要在本地起服务打开浏览器就能用不需要折腾前端工程。下面是最小可用的接口代码import io import torch import torchvision.transforms as T from PIL import Image from flask import Flask, request, jsonify app Flask(__name__) device torch.device(cpu) # 演示机器优先 CPU别赌 GPU # map_location 强制映射到 CPU避免 GPU 训练的权重在 CPU 机器上加载失败 checkpoint torch.load(checkpoints/mnist_cnn.pt, map_locationdevice) model DigitCNN() model.load_state_dict(checkpoint[model_state]) model.eval() transform T.Compose([ T.Resize((28, 28)), T.ToTensor(), T.Normalize((0.1307,), (0.3081,)) ]) app.route(/predict, methods[POST]) def predict(): file request.files.get(image) if file is None: return jsonify({error: missing image}), 400 try: img Image.open(io.BytesIO(file.read())).convert(L) except Exception: return jsonify({error: invalid image}), 400 tensor transform(img).unsqueeze(0) # 加 batch 维度: (1,28,28) with torch.no_grad(): probs torch.softmax(model(tensor), dim1)[0] digit int(probs.argmax().item()) confidence float(probs.max().item()) top2 [int(i) for i in probs.topk(2).indices.tolist()] return jsonify({ digit: digit, confidence: confidence, top2: top2 }) if __name__ __main__: app.run(host0.0.0.0, port5000)这个接口做了三件关键的事。第一map_locationcpu保证不管训练时用的什么设备加载时都走 CPU否则在只有 CPU 的答辩电脑上会直接报 CUDA 相关错误。第二convert(L)把上传图片强制转成灰度因为 MNIST 模型输入是单通道如果来一张 PNG 或彩色照片RGB 三通道会让 shape 对不上。第三返回里带了confidence和top2这是给第 6 章的“不确定提示”预留的。关于Image.open有一个隐蔽的坑它并不会真正读取图片像素只是打开一个文件句柄直到.convert()时才实际解码。所以如果上传的不是合法图片Image.open不会立刻报错而是在 convert 时抛异常。接口里用 try/except 包住解码过程返回 400 而不是让服务直接崩掉这个对演示稳定性很重要。4.2 前端画板与上传两种输入方式从 canvas 到 base64接口有了还需要一个能让使用者“手写”的前端。最稳妥的方案是 HTML canvas 画板 文件上传两条路都做画板用于现场演示上传用于验证真实图片。先看画板的核心部分canvas idboard width280 height280/canvas button idclear清空/button button idrecognize识别/button script const canvas document.getElementById(board); const ctx canvas.getContext(2d); let painting false; // 黑底白字跟 MNIST 的视觉习惯一致 ctx.fillStyle black; ctx.fillRect(0, 0, 280, 280); ctx.strokeStyle white; ctx.lineWidth 18; ctx.lineCap round; // 完整做法是监听三个事件 // mousedown 时 painting true 并 beginPath // mousemove 时 lineTo 再 stroke // mouseup 时 painting false canvas.addEventListener(mousedown, () { painting true; }); canvas.addEventListener(mousemove, drawLine); canvas.addEventListener(mouseup, () { painting false; }); function drawLine(e) { if (!painting) return; const rect canvas.getBoundingClientRect(); const x e.clientX - rect.left; const y e.clientY - rect.top; ctx.lineTo(x, y); ctx.stroke(); } document.getElementById(recognize).onclick async () { const dataUrl canvas.toDataURL(image/png); const blob await (await fetch(dataUrl)).blob(); const form new FormData(); form.append(image, blob, digit.png); const res await fetch(/predict, { method: POST, body: form }); const data await res.json(); alert(识别结果${data.digit}置信度 ${(data.confidence * 100).toFixed(1)}%); }; /script画板设计里最容易忽略的问题是canvas 尺寸是 280x280但模型需要的是 28x28。直接缩小后笔画会变得很细很碎很多数字会认错。我一般会加一个“居中裁剪”的预处理先用getImageData找到画板里所有非零像素的包围盒把包围盒裁出来扩展 10% 的留白再绘制到 28x28 的黑色画布中央。这一步对真实手写识别的提升比换模型还大因为 MNIST 训练数据的数字基本都是居中的而你在 280x280 画板上写字很难精确居中。文件上传路径更简单一个input typefile接同一个/predict接口就行。要注意的是手机拍的照片通常不止一个数字而且有背景干扰这个题目一般不做多数字分割所以上传功能定位成“验证单数字图片”不要承诺能识别整张纸。另外如果前端页面和 Flask 服务不是同一个端口浏览器跨域会拦截请求需要给 Flask 加 CORS 头省事的做法是直接把 HTML 放进 Flask 的 templates 目录由 Flask 渲染这样同源就没有跨域问题。4.3 接口联调与稳定性curl 测试与并发注意前后端写完后先用命令行把接口测通再打开浏览器能省掉大量联调时间。测试命令curl -X POST -F imagetest_digit.png http://127.0.0.1:5000/predict正常返回应该是类似{digit: 3, confidence: 0.98, top2: [3, 8]}的 JSON。如果返回 404 或 500先看 Flask 控制台输出的异常栈大部分时候是图片解码失败或预处理尺寸不对。关于并发有一个容易被忽略的事实Flask 自带的开发服务器是单进程的模型推理本身在 CPU 上单张约几十毫秒演示时如果连续快速点击识别按钮请求会排队界面上像是卡住了。常见做法是给按钮加一个禁用状态请求发出后立刻把按钮置灰收到响应再恢复。至于换 gunicorn 或加线程锁毕设演示没有这个必要反而徒增部署复杂度。另外模型加载要在app.run()之前完成放在全局变量里不要在每次请求里去torch.load——每请求加载一次权重不仅慢而且内存会迅速吃紧这是我见过最典型的部署错误。5. 手写数字识别毕设避坑训练不收敛、精度上不去与答辩追问5.1 训练 Loss 不下降学习率与权重初始化是第一怀疑对象现象训练了好几轮loss 一直稳定在 2.3 左右。懂的人知道log(10) 约等于 2.3026这正是“模型在十个类别里随机瞎猜”的损失值说明它什么都没学到。原因排序第一数据预处理出问题最常见是 Normalize 根本没生效比如 transform 写错输入还是 0-255 的原始像素或者是标签和图像错位第二学习率不合适Adam 默认 1e-3 在 MNIST 上基本不会出错但如果你用了更大的学习率比如 1e-1loss 会在高位震荡第三权重初始化极端值不过用 PyTorch 默认初始化时这几乎不用排查。解决先打印一个 batch 的 images 和 labels肉眼确认像素值在 0-1 附近、均值接近 0.13标签范围是 0-9。然后把学习率降到 1e-3 或 1e-4 重跑。如果还不降就换一个极小数据集比如 500 张去过拟合如果小数据集能降到接近 0说明数据链路没问题问题在大训练集上的配置如果小数据集也降不下去问题在模型代码本身优先检查卷积输出的形状和全连接层的输入尺寸是否匹配。5.2 精度卡在 95% 上不去检查预处理与验证集划分现象训练集精度一路涨到 99%验证集或测试集却卡在 95% 左右怎么调参都上不去。原因最常见的是没有验证集做早停。你用了全部 60000 张训练、看的是最后一个 epoch 的模型那个模型已经过拟合正好在测试集上表现回落。其次可能是训练和测试的预处理不一致训练时用了数据增强测试时也应该用同一套 base transform而不是把增强也带上反过来训练时忘了归一化、测试时才归一化也会导致分布错位但这种情况精度会掉得更离谱。解决切出 5000 张验证集在训练循环里每轮记录 val_acc保存 val_acc 最高那个 epoch 的权重而不是最后一轮。如果验证集精度正常但测试集低再检查测试集评估时是不是忘了model.eval()或者测试集被意外 shuffle 过。这两个问题症状相似但排查路径完全不同。MNIST 上 95% 意味着大约 500 张测试图被认错这个量级不是网络容量问题而是流程问题别急着换大模型。5.3 自己写的数字识别错真实手写与 MNIST 的分布差异现象测试集 99%用画板写一个好好的“5”识别成“6”写“7”识别成“1”。看起来是模型太笨其实是数据分布差异。原因MNIST 的训练样本是 28x28、数字居中且笔画均匀的灰度图。浏览器画板出来的是抗锯齿笔画线宽只有几像素而且你的字大概率偏左或偏上缩放成 28x28 后和训练样本完全不像。这是“真实手写”和“MNIST 手写”之间的领域差异任何模型都扛不住盲目输入。解决不要在模型上死磕在前端预处理上做三件事。第一找到画板非零像素的包围盒把数字裁出来第二加一圈留白后缩放到 20x20再贴到 28x28 黑色画布正中央第三如果笔画太细用 OpenCV 的膨胀操作加粗一下让线宽接近 MNIST 的笔画粗细。按这个流程大部分识别错误能直接消除。这也是我为什么坚持让项目带“额外样本”——拿 30 张自己写的数字跑一遍你就知道模型真实水平而不是测试集分数。5.4 答辩现场模型加载失败路径、版本与设备不匹配的排查现象在自己电脑上一切正常拷到答辩机器的 U 盘里运行python app.py后在加载模型时报错或者打开页面后识别接口直接 500。原因排序第一checkpoint 里保存的是model.state_dict()但加载脚本里没有先定义DigitCNN类或者类名不一致导致反序列化失败第二训练时在 GPU 上保存的权重加载时map_location没有指定 CPU第三项目路径带中文或空格torch 加载文件时在某些环境下会出错第四Python 和 torch 版本不一致比如本机 Python 3.10 加 torch 2.0答辩机器是 Python 3.8 加 torch 1.12极少数情况下会有兼容问题。解决用第 3 章那种“带元信息的 checkpoint”写法加载时先实例化自己的模型类再load_state_dict完全不依赖序列化时的类路径torch.load一律加map_locationcpu项目文件夹命名为mnist_demo这种纯英文路径放到桌面或用户目录下不要放进带中文的文件夹。最后答辩前一晚一定要做一次“从零到一”演练换一台没有装任何 Python 包的机器按 README 从 conda 环境装起跑通一次完整流程。这个过程会暴露所有你以为已经解决了的环境依赖问题。6. 从 99% 到更稳置信度阈值与 Top-2 提示的实战技巧6.1 混淆矩阵与错误样本怎么看模型到底错在哪98% 和 99% 的差别在 MNIST 测试集上就是几十张图这时候看总精度没有意义要看混淆矩阵。最常见的混淆对是 4 和 9、7 和 9、3 和 8原因是这些数字的局部结构相似。把你预测错的样本用 matplotlib 画成一张网格图哪一类错得多、错的是什么风格一眼就能看出来。如果错误样本集中在“笔画断裂”或“倾斜过大”说明预处理增强不够如果集中在某个特定类别可以考虑给这个类别多采集点变体样本或者调整类别权重。这些分析讲出来答辩的深度立刻不一样。6.2 置信度阈值与 Top-2 提示不增加训练量的兜底策略这个技巧不训练任何新模型但对答辩演示的提升立竿见影。softmax 输出的概率本身就带有不确定性信息一个样本如果置信度只有 0.55与其硬报一个可能错的答案不如在界面上提示“不太确定可能是 5也可能是 6”。实现只需要在接口返回里加一个判断digit probs.argmax().item() confidence probs.max().item() top2 [int(i) for i in probs.topk(2).indices.tolist()] if confidence 0.7: result f不太确定可能是 {digit}也可能是 {top2[1]} else: result f识别为 {digit}置信度 {confidence:.2f}这个阈值 0.7 是我自己常用的默认值你可以用验证集统计一下把置信度低于 0.7 的样本全列出来看看里面错误率是不是显著高于整体。如果是说明阈值设置合理如果 0.7 太低拦不住错误就提高到 0.8 甚至 0.85。这个做法在工程上叫拒绝识别很成熟也是答辩时能拿出来讲的加分点你证明了模型不仅能给出答案还知道自己什么时候不确定。我现在的习惯是交付这类项目前一定跑一遍“三张图测试”一张标准的 MNIST 测试图一张自己用画板写、居中良好的数字一张故意写歪、笔画很细的数字。三张图的识别结果分别是什么置信度分别是多少心里有数才敢拿到答辩现场。否则演示时随机写一个数字翻车前面讲得再好都很被动。希望这套从数据到部署的流程能帮你把这个经典题目做成真正意义上的“系统”而不是一个只有 ipynb 的模型实验。希望帮到你。本文还有配套的精品资源点击获取
返回列表