ARTICLE DETAIL

资讯详情

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

PyTorch+MNIST+CNN实战:从数据下载到模型部署全流程踩坑记录

PyTorch+MNIST+CNN实战:从数据下载到模型部署全流程踩坑记录 要不是昨天有人在群里发了一张 torchvision 下载 MNIST 数据集报 404 的截图我差点忘了自己当年入坑 PyTorch 时也被这一下卡了整整一个晚上。想学 CNN 却连最经典的手写数字分类数据集都拉不下来这几乎是每个初学者都会经历的一道“开胃败仗”。这篇东西不是官方文档的复述而是我实际从零跑通 PyTorch MNIST CNN 全流程的记录。包括环境怎么搭、数据加载器怎么写、模型结构为什么这样堆、训练循环里哪些细节不处理必踩坑以及最后如何把模型保存下来、导出 ONNX、用它识别真实手写图片。我会把自己踩过的 404、CUDA 不可用、黑白翻转这些坑一一说清楚把你最可能卡住的位置提前点上灯。1. 为什么MNIST是CNN入门路上绕不开的那个数据集不少人会问都 2024 年了MNIST 这种 28×28 的老古董还有什么好学的我的看法是正因为简单它才是最好的“试金石”。MNIST 由 6 万张训练图片和 1 万张测试图片组成每张都是 28×28 的灰度图内容是 0 到 9 的手写数字。它的特点非常鲜明单通道、小尺寸、类别均衡、背景干净。这意味着你不需要多贵重的 GPU普通 CPU 都能在几分钟内把模型训练到 99% 以上。对初学者来说这种“付出有回报、调试有反馈”的节奏感太重要了。从另一个角度讲MNIST 又足够“像样”。手写数字存在笔画粗细不均、位置偏移、形变、断笔等现象这些特征和真实图像识别任务面临的问题是同一类问题只是程度更轻。你在 MNIST 上学会的卷积、池化、数据增强、过拟合分析、训练评估流程换到 CIFAR-10、ImageNet 时依然成立只是数据和模型规模变大而已。还有个很容易被忽略的价值MNIST 的错误可视化非常直观。训练集上预测错了把图片画出来肉眼看一眼就知道是模型犯糊涂还是标注本身就有问题。这种“即时纠错”的体验对建立直觉极有帮助不是每个数据集都能给你。所以别嫌它“太简单”。我见过一些新同学一上手就啃 Transformer、Diffusion结果连梯度都不回传愁眉苦脸一整天。先把 MNIST 上的 CNN 流程吃透后面的复杂网络才谈得上有章法。2. 环境搭建踩坑实录版本匹配、CUDA、WSL与MNIST下载404环境搭建是最劝退新手的一步。这里啰嗦几句我踩过的坑几乎都是在这里踩的。2.1 安装策略conda 还是 pip日常开发怎么选我的习惯是先用 Anaconda 建一个独立环境避免把系统的 Python 环境搅乱。操作非常简单conda create -n mnist python3.10 conda activate mnist接着安装 PyTorch。官网首页会给出对应 CUDA 版本的安装命令这个必须自己去查因为 PyTorch 版本、CUDA 版本、Python 版本三者之间有兼容关系。选 pip 还是 conda我个人的建议是能用 conda 就用 conda尤其在 Windows 上。conda 对 MKL、OpenMP 这类底层依赖的处理更省心pip 偶尔会遇到 PyTorch 装好了但 import 时动态链接库报错的情况。这里最容易被忽略的是 CPU 版和 GPU 版的区别。如果电脑没有 NVIDIA 显卡直接装 CPU 版就够了如果有显卡务必看清楚安装命令里是 cu118、cu121 还是 cu124 这种后缀它代表配套的 CUDA 版本。不要盲目装最新的 CUDAPyTorch 官方的支持列表是你该相信的不是显卡驱动面板里显示的数字。安装完一个习惯性自检import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0) if torch.cuda.is_available() else CPU)如果torch.cuda.is_available()返回 False是大多数环境问题集中爆发的信号。原因可能包括装成了 CPU 版、NVIDIA 驱动版本太老、系统 CUDA 版本与 PyTorch 不匹配、在 WSL 里没启用 GPU 支持等。2.2 在 WSL/Linux 环境里配置 GPU 支持现在很多同学在 WSL 里跑实验。WSL 2 本身支持 CUDA前提是 Windows 侧已经装好 NVIDIA 显卡驱动并且 WSL 里的 PyTorch 必须装 CUDA 版。我遇到的一个典型情况是在 WSL 里nvidia-smi能正常显示但 PyTorch 就是检测不到 CUDA。原因往往是torch.cuda.is_available()检测的是 PyTorch 编译时链接的 CUDA runtime而不是nvidia-smi里的驱动版本。解决办法是在 WSL 的 conda/pip 环境里重新安装对应 CUDA 版本的 PyTorch而不是依赖系统级 CUDA。还有个容易踩的坑Windows 和 WSL 的文件系统互访。如果把项目放在/mnt/c/...路径下训练时读取小文件可能很慢。MNIST 这种小数据集还好如果换到大数据集建议把代码和数据都放在 WSL 的 ext4 目录里速度和稳定性会好很多。2.3 torchvision 下载 MNIST 报 404这里有一个手工落地的方案很多同学是在这一步崩溃的执行datasets.MNIST(root./data, downloadTrue)后终端出现HTTP Error 404: Not Found数据集下载失败。这个问题的根源通常是下载源访问不顺畅。MNIST 数据集文件本身是公开的、体积也就几 MB 到十几 MB但下载源在某些网络环境下并不稳定。我不建议反复重试更靠谱的是手工下载这 4 个文件train-images-idx3-ubyte.gztrain-labels-idx1-ubyte.gzt10k-images-idx3-ubyte.gzt10k-labels-idx1-ubyte.gz下载完成后把它们放到./data/MNIST/raw/目录下注意目录层级必须匹配。然后把代码里downloadTrue改成downloadFalse再运行一次。torchvision 检测到本地已经有原始文件就会自动解压并生成格式化的数据。这个手工方案的原理其实很简单datasets.MNIST的第一步是从远端拉取原始 gzip 文件如果本地 raw 目录已经有这些文件它就不管网络了。从教育角度讲手工下载文件也让你更清楚数据从哪来必要时还能编写脚本自行校验。另外补充一点如果运行环境本身访问下载源比较慢即使不报 404也可能卡很久。这时候可以设置一个超时时间仔细观察报错信息不要盲目等。2.4 训练时遇到“CUDA out of memory”怎么办MNIST 数据集很玩具基本不会 OOM除非你把 batch size 调到几千。但很多人用同一个 torch 环境跑大网络时就会撞上。养成好习惯数据、模型、损失函数全部to(device)batch size 不要盲目翻倍。如果真的遇到显存不足先把batch_size减半或者减少 DataLoader 的num_workers这招在 MNIST 阶段能覆盖大多数情况。3. 数据管线Dataset、DataLoader和归一化里的小心计模型写得再漂亮数据管线没做好也白搭。MNIST 的数据读取用 torchvision 几行就能搞定但背后的几个细节值得掰开揉碎讲。先看一段常用写法from torch.utils.data import DataLoader from torchvision import datasets, transforms transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_set datasets.MNIST(root./data, trainTrue, transformtransform, downloadFalse) test_set datasets.MNIST(root./data, trainFalse, transformtransform, downloadFalse) train_loader DataLoader(train_set, batch_size128, shuffleTrue, num_workers4, pin_memoryTrue) test_loader DataLoader(test_set, batch_size256, shuffleFalse, num_workers4, pin_memoryTrue)3.1 ToTensor 不只是“转张量”这么简单transforms.ToTensor()完成两件事把 PIL Image 或 numpy 数组H×W×C转成 PyTorch 张量并在最后增加一个通道维度C×H×W同时把像素值从 0~255 线性缩放到 0~1。很多第一次写代码的人会自己手动做x torch.tensor(np.array(img)).float()结果训练直接崩一脸。原因很典型通道顺序不对或者 dtype 不是浮点型、范围没有归一化。ToTensor帮你把这些琐碎细节一次性处理干净能少踩很多坑。MNIST 原始图片是单通道灰度图所以ToTensor()之后张量形状是[1, 28, 28]第一个 1 就是通道维。如果后续模型第一层Conv2d的in_channels1正好对上。3.2 归一化为什么非做不可灰度图转成 0~1 范围已经不错了但还不够。Normalize((0.1307,), (0.3081,))里的两个数字是 MNIST 全体训练集像素的均值和标准差。为什么要额外归一化因为神经网络对输入的尺度很敏感。输入特征的值域如果不一致早期梯度更新的方向就会有很大波动训练要么不稳定要么收敛很慢。把数据变成“零均值、单位方差”之后每个像素的特征分布更接近优化器迭代起来也快得多。这两组数字并不是拍脑袋定的而是统计出来的所以直接沿用就好。这个参数等模型训练完、部署新图片时还要再一次用到。很多人模型训练没问题但预测单张图时忘了套同样的归一化结果准确率莫名其妙掉到 20% 以下这点后面还会再提。3.3 数据增强MNIST 上简单做就行别上头有人刚学到数据增强就恨不得把旋转、平移、缩放、加噪声全部堆上去。但 MNIST 做增强要克制。数字 6 和 9 只要稍加旋转就容易混淆人眼都容易认错过度旋转反而让测试集性能下降。如果确实想加我推荐两种轻量方式RandomAffine加一个很小的角度范围比如正负 10 度外加 10% 左右的平移RandomErasing随机擦除一小块模拟笔画不完整的情况。在 MNIST 上做增强的问题在于测试集和训练集分布差异不大增强带来的泛化收益有限反而增加了训练时间。我的经验是先跑一次不加增强的基线精确率大约 99% 左右再决定是否加增强否则你很难判断提升究竟是增强带来的还是代码修改带来的。3.4 DataLoader 的 tiny 细节batch_size在没有显卡的情况下不宜太大否则计算慢。CPU 上我建议batch_size64或128GPU 上256也毫无压力。shuffle只对训练集有意义测试集不要 shuffle否则你后续打印一张图对一张标签时顺序全乱掉排查问题会很难受。num_workers在 Linux 上可以设成 4 或者 8但在 Windows 上经常会报BrokenPipeError或者DataLoader worker (pid...) is killed。这是 Windows 多进程的经典问题省事方案是设成num_workers0。虽然数据读取慢一点但能避免一个巨坑。我自己的经验是MNIST 数据集太小num_workers0的额外开销完全在可接受范围没必要为了炫技给自己找麻烦。pin_memoryTrue在 GPU 训练时可以把数据固定到页锁定内存里减少主机到 GPU 的拷贝耗时。但它必须配合devicecuda才有意义纯 CPU 训练时开着不会有坏处也不用指望有多大提升。4. 搭建CNN模型从LeNet思路到现代小改动MNIST 的 CNN 网络结构并不复杂经典 LeNet-5 已经能做到很好的效果。不过 LeNet 是上世纪 90 年代的设计很多写法在现代 PyTorch 里可以做点小改进比如加上 BatchNorm 和 Dropout。我常用的一个结构是这个import torch import torch.nn as nn class MNISTCNN(nn.Module): def __init__(self): super().__init__() self.features nn.Sequential( nn.Conv2d(1, 32, kernel_size3, padding1), nn.BatchNorm2d(32), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size3, padding1), nn.BatchNorm2d(64), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(64 * 7 * 7, 128), nn.ReLU(inplaceTrue), nn.Dropout(0.2), nn.Linear(128, 10), ) def forward(self, x): return self.classifier(self.features(x))4.1 每一层为什么这么放先看输入输出尺寸的变化。输入是[1, 28, 28]。第一次卷积kernel_size3, padding1会让尺寸保持不变仍然是28×28因为(28 - 3 2*1) / 1 1 28。接着MaxPool2d(2)把宽高减半变成14×14。第二次卷积同理保持14×14再池化变成7×7。此时通道数是 64所以展平后是64*7*7 3136。用 padding1 的好处是让特征图尺寸成倍数缩小计算形状很直观。如果不加 padding第一层卷积后变成26×26池化后13×13第二次卷积后11×11池化后5×5展平尺寸是64*5*5 1600。两者都能用但我建议初学者用 padding1 的写法减少形状推算时的出错概率。两个卷积块就够了吗对 MNIST 来说两个卷积层已经能捕捉到笔画、边缘、闭环结构等关键特征。继续堆更多卷积层精度的提升非常有限还会让训练更慢、过拟合风险更高。训练 CNN 不是层数越多越好而是够用就好。4.2 BatchNorm 与 Dropout 的位置奥妙BatchNorm2d放在卷积之后、ReLU 之前已经成为一种主流做法。它的作用是把当前 batch 的分布拉回均值为 0、方差为 1 的状态缓解网络内部协变量偏移问题。但 BatchNorm 在训练和推理时的行为是不一样的。训练时它统计当前 batch 的均值和方差推理时它使用训练过程中累计得到的全局均值和方差。这就是为什么后面训练循环里model.train()和model.eval()必须成对切换否则会出现很诡异的现象训练时 loss 正常下降测试时结果一塌糊涂。Dropout(0.2)放在全连接层之前的含义是每次前向传播有 20% 的神经元输出被随机置零迫使网络不依赖个别神经元提高泛化能力。Dropout 同样只在训练时生效在eval()模式下它会自动关闭。这两个层倒不是说对 MNIST 提升巨大但它们教会你的知识点会沿用到所有后续网络里。4.3 为什么不急着上残差网络现在很多人一上来就 ResNet这没必要。ResNet 处理的是深层的梯度退化问题它需要足够深的网络才能体现优势。在 MNIST 上强行套 ResNet反而网络过重训练慢而且要处理输入尺寸、padding、stride 等问题对新手极不友好。我的建议是先把上面这个轻量 CNN 从零写明白弄懂 forward 里每一步的形状变化、每个模块的作用再逐步引入更重的结构。真正的高手不是只会调包而是对组合结构有直觉。5. 训练循环CrossEntropy、Adam与99%背后的工程细节模型定义好接下来是训练。MNIST 的训练循环说简单也简单说难也有几个细节会卡住人。完整代码参考这段import torch import torch.optim as optim device torch.device(cuda if torch.cuda.is_available() else cpu) model MNISTCNN().to(device) criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr1e-3) for epoch in range(5): model.train() total_loss 0.0 correct 0 for x, y in train_loader: x, y x.to(device), y.to(device) optimizer.zero_grad() out model(x) loss criterion(out, y) loss.backward() optimizer.step() total_loss loss.item() * x.size(0) pred out.argmax(dim1) correct (pred y).sum().item() train_acc correct / len(train_loader.dataset) print(fepoch {epoch1}, train loss: {total_loss / len(train_loader.dataset):.4f}, train acc: {train_acc:.4f}) model.eval() test_correct 0 with torch.no_grad(): for x, y in test_loader: x, y x.to(device), y.to(device) out model(x) pred out.argmax(dim1) test_correct (pred y).sum().item() test_acc test_correct / len(test_loader.dataset) print(ftest acc: {test_acc:.4f})5.1 损失函数为什么是 CrossEntropyLoss 而不是 MSELoss分类任务的标准选择是交叉熵损失对应 torch 里的nn.CrossEntropyLoss。这里有个细节必须说清楚CrossEntropyLoss内部已经包含 softmax 操作。也就是说它接收的是模型输出的原始 logits而不是经过 softmax 的概率值。很多人第一次手写模型时会在nn.Linear之后再加一个nn.Softmax然后把自己坑了。原因是logits 在进入 CE loss 前还要再过一次 log_softmax你手动的 softmax 反而让梯度在数值上变得不稳定而且和网络最后的概率标签对不上。把nn.CrossEntropyLoss理解成“log_softmax NLLLoss”的合并体就够了。模型最后一层不用 softmax它在训练时由 loss 函数内部处理在推理时你用out.argmax(dim1)拿到预测标签即可。5.2 优化器选 Adam 还是 SGD初学者我推荐 Adam学习率设1e-3基本不需要太精细的调参就能快速收敛。原因在于 Adam 自带自适应学习率对梯度尺度不敏感比较“皮实”。但我也得坦白说真正调参经验更丰富之后SGD with Momentum 往往能刷出比 Adam 更高的测试精度。SGD 对学习率更加敏感需要更精细的 schedule但它收敛到的最小值通常更平缓泛化性更好。MNIST 任务上这个差异很小没必要为了 0.1 个百分点折腾太久。经验值供参考MNIST 上 Adam lr1e-3通常 3 个 epoch 训练集就有 99.5% 以上5 个 epoch 测试集能达到 99.2% 到 99.5%。再往高刷的边际成本远大于收益。5.3 model.train() 与 model.eval() 的切换不能省略很多初学者写完循环不切model.train()和model.eval()之前 4.2 我讲过 BatchNorm 和 Dropout 的存在它们对 train 和 eval 两种状态反应完全不同。model.train()是在告诉模型现在是训练阶段BatchNorm 使用当前 batch 统计量Dropout 随机失活。model.eval()则进入推理模式BatchNorm 使用训练期间保存的全局统计量Dropout 关闭同时nn.Flatten之外的任何可能影响结果的结构都要稳定下来。如果把这段切换忘了测试时模型内部某些层的行为仍然保持训练模式输出结果会出现波动准确率看起来忽上忽下。这是新手最容易忽略但又最影响判断的坑。5.4no_grad的作用测试阶段用with torch.no_grad():包裹起来。它的意思是告诉 PyTorch 不需要计算梯度。如果你不加这层保护测试前向传播仍然会构建计算图、暂存中间变量不但占内存还浪费时间。在 MNIST 这种小任务上影响不明显但换成大一点的模型比如后续做 CIFAR-10 时这个习惯必须提前养成。5.5 训练过程的实时监控不是可选项我强烈建议每个 epoch 打印训练集 loss、训练集准确率和测试集准确率。别只盯测试准确率因为训练集 loss 的变化能告诉你模型是否还在收敛是否已经开始过拟合。一个经典信号是训练集准确率不断逼近 100%但测试集准确率停滞甚至下降。这是过拟合的警报。这时候可以做的事包括增加 Dropout、加入轻量数据增强、减少训练 epoch而不是继续盲刷训练轮数。6. “准确率挺高”是好事但MNIST会骗人过拟合与数据分布很多人看到测试集 99.4% 就觉得自己“搞定了”。实际上 MNIST 太简单准确率非常容易被刷上去也正因为这样它非常适合用来分析过拟合和分布漂移。6.1 训练几个 epoch 就过拟合了MNIST 的 6 万张训练图片相比模型参数量并不算多。当模型训练到第 4、5 个 epoch 时训练集的准确率往往会达到 99.9% 以上而测试集却停留在 99.3% 左右。这条“缝”就是过拟合。我习惯的做法是这样的把训练集每个类的准确率单独算一遍再看测试集的混淆矩阵重点关注到底哪些数字互相混淆。比如常见的 4 和 9、3 和 8、7 和 2这种混淆在笔画结构上很合理如果连人眼都觉得难分那模型的错误就有可解释性。6.2 测试集不能拿来反复调参这里有一条我在学习和带人时反复强调的规则测试集只能你最终用一次用来报告效果不能作为日常调参的指标。因为一旦你根据测试集性能反复修改超参数、模型结构测试集的信息就“污染”了模型选择过程最终报告的数字会失真。正确做法是把训练集再切出一块验证集或者用交叉验证调参只看验证集。MNIST 官方给了测试集但很多人不自觉地把测试集当成验证集用。虽然 MNIST 太简单这个做法坏处不大但坏习惯一旦养成后面做大项目时一定会吃亏。6.3 一张真实手写图片就把模型打回原形MNIST 图片的背景是黑色、数字是白色。现实里拿手机拍一张白纸黑字的手写数字导入模型前如果只是简单地 resize 到 28×28、转成灰度图大概率识别效果惨不忍睹。原因有两层颜色极性是反的。MNIST 训练数据大多是黑底白字白底黑字输入相当于把图像颜色反转了模型看到的特征分布彻底变了。数字不一定居中笔画粗细也相差很大。MNIST 的预处理把数字大致居中但现实图片随手一拍的偏移和缩放网络没见过就很容易误判。解决方案是在预处理里多做几步把图片转成灰度、根据阈值做二值化、裁掉过于冗余的边缘、缩放到 20×20 再放到 28×28 画布中央。这一步说白了是尽量把真实图片处理成和 MNIST 训练分布一致的样子。很多人在模型上折腾半天结果问题出在预处理上这是最尴尬的翻车现场。6.4 保存错例比保存准确率数字更有用训练完后我会专门做一个步骤把测试集里预测错误的那批(image, label, pred)三元组保存成图片按类别分开存放。这样我能直观看到模型在哪些数字上犯迷糊。这种做法比贴一个“test acc: 0.9930”有用得多。因为准确率只是一个标量错例却是可解释的证据。哪怕你没时间画混淆矩阵抽几十张错例图看一遍也能快速判断问题是数据预处理、模型容量还是训练策略导致的。这个习惯从 MNIST 阶段养成会让你后面处理真实项目时受益很多。7. 保存模型、导出ONNX以及拿真实手写图片推理到这里模型已经训练好了但博文不能止步于训练后面还有几张“实战底牌”。7.1 正确保存和加载方式PyTorch 保存模型有几种写法我建议尽量保存state_dict也就是模型权重字典而不是整个模型对象torch.save(model.state_dict(), mnist_cnn.pth)加载的时候需要先定义同样的网络结构再加载权重model MNISTCNN() model.load_state_dict(torch.load(mnist_cnn.pth, map_locationcpu)) model.eval()用map_locationcpu的好处是模型如果在 GPU 上训练加载到没有 GPU 的机器上也不会报设备不匹配的错误。保存整个 model 对象虽然省事但一旦网络代码变化或者依赖版本升级反序列化很可能出问题实际部署时非常脆弱。7.2 导出为 ONNX 的流程导出 ONNX 是为了把模型从 PyTorch 生态中解脱出来方便部署到移动端、服务端或者嵌入到推理引擎里。MNIST 这个例子导出 ONNX 非常简单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}} )dynamic_axes设置成动态 batch 后推理时可以一次传入多张图片不至于每次只能处理单张。导出完成后可以用onnxruntime验证一遍输入输出的形状确保和预期一致。还有一个细节ONNX 导出时模型要在eval()模式。如果在train()模式下导出BatchNorm 和 Dropout 的行为仍然停留在训练状态导出后的模型在推理端结果会不稳定。7.3 用真实手写图片做推理这里我用一个非常简化的单张图片推理示例前提是图片已经过灰度化、二值化、裁剪缩放处理最终得到 28×28 灰度张量import cv2 import numpy as np import torch from torchvision import transforms def load_and_preprocess(img_path): img cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) img cv2.resize(img, (28, 28), interpolationcv2.INTER_AREA) # 如果检测到背景偏黑、前景偏白可以先保持原样 # 如果背景是白色需要反转像素值让数字变成白色 if np.mean(img) 127: img 255 - img img img.astype(np.float32) / 255.0 img (img - 0.1307) / 0.3081 img torch.from_numpy(img).unsqueeze(0).unsqueeze(0) # [1,1,28,28] return imgnp.mean(img) 127这一步是我的常用判断如果整图平均亮度很高说明大概率是白底黑字直接翻转像素值让数字变成高亮像素而不是黑底白字。这个自动判断不是万无一失但对大多数真实图片足够有效。推理部分就是常规操作with torch.no_grad(): pred model(img).argmax(dim1).item() print(pred)这里最容易出问题的地方是张量形状。灰度图的形状是[H, W]加两个维度后变成[1, 1, H, W]如果写成img.unsqueeze(0)少了通道维模型会直接报维度错误或者in_channels1与输入通道数 3 对不上。还有同学习惯用 PIL 的Image.open却忘了convert(L)把三通道 RGB 图直接喂进单通道模型也会翻车。7.4 把整个推理流程封装成一个类后续省力等模型稳定之后我就不再写裸函数了而是把它封装成一个小类方便在 API 服务和命令行工具里复用。核心是三个方法preprocess、predict、predict_batch。类的好处是初始化时加载模型一次推理时不重复加载服务响应速度快很多。封装的时候注意一个细节模型加载后必须调用model.eval()同时保持 forward 只做前向传播不在类内部改任何参数。这样你在调试和部署时行为和结果都可复现不会出现一个接口今天识别成功明天识别结果不同的奇怪现象。我在实际使用中的体会是MNIST 这个项目虽然小但它几乎覆盖了一个深度学习完整流程里所有的关键节点环境配置、数据处理、模型设计、训练评估、导出部署、真实场景适配。把这套流程跑顺了后面的路会走得踏实很多。最后再分享一个小技巧训练完成后把测试集里预测错误的图片全部保存到一个单独目录每隔一段时间拿出来看看。你会发现模型真正学到的和你以为它学到的经常是两个完全不同的东西。
返回列表