ARTICLE DETAIL

资讯详情

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

MNIST手写数字识别:从传统机器学习到CNN的完整实战

MNIST手写数字识别:从传统机器学习到CNN的完整实战 简介基于机器学习方法的MNIST手写数字识别项目使用Python 3.6分别实现SVM、决策树、KNN、朴素贝叶斯四种算法并在同一数据集上比较识别准确率。压缩包共19个文件、约11.04MB其中含4个Python源码、标准MNIST训练/测试图像与标签文件、各算法结果及准确率对比图、训练模型文件、决策树结构文档及说明文档代码按Code、Dataset、res目录清晰组织。截至目前已有539人学习或下载。这套资源对计算机、人工智能等相关专业学生尤为实用代码经过完整测试、运行成功既可作为毕业设计、课程设计或项目初期演示的参考也可在此基础上自由替换算法或调参是理解手写数字识别并横向对比不同分类器效果的优质资料。1. 为什么 MNIST 至今仍是机器学习入门绕不开的试金石这两年做大模型的朋友经常开玩笑说“MNIST 太老了老到新手都不屑于跑”但真到自己带团队、带实习生我第一个布置的任务还是 MNIST。原因很简单它足够小、足够干净、标签明确、训练快能把“机器学习应用流程”从头到尾完整走一遍——从读数据、做特征、选模型、调参到评估每一步都有直接的反馈。而“机器学习检测”这类听起来高阶的任务落到工程上做的其实还是 MNIST 这套基本功的放大版。很多人初次接触“MNIST 手写数字识别”时以为必须上深度学习才行甚至直接去找 PyTorch 的现成源码。实际上传统机器学习算法在这份 28x28 灰度图上同样能跑到 97% 上下而那一套特征工程加上模型融合的思路才是“机器学习算法”和“计算机视觉”之间那条最清晰的界线。无论你是刚看完周志华《机器学习》准备复现课后题还是期末前想拿一个能写进简历的完整项目这个标题里提到的“源代码 文档说明 数据集”三件套本身就是一份标准的入门工程模板。下面我用自己常跑的流程把这条路线完整拆开。2. 先认识数据集MNIST 的字段、读取方式和传统机器学习基线2.1 数据集结构70000 张图是怎么被划分的MNIST 由 Yann LeCun 等人整理训练集 60000 张、测试集 10000 张每张图片 28x28 像素灰度范围 0 到 255。数据集的原始格式是 IDX 二进制不是常见的 PNG 或 JPG这意味着你不能直接cv2.imread必须自己解析文件头。很多新手第一次在这上面“翻车”就是因为直接用读图片的方式去处理 IDX 文件读出来一堆乱码。整个文件结构分为四个部分训练图像、训练标签、测试图像、测试标签。图像文件中前 16 个字节是魔数和维度信息标签文件前 8 个字节是文件头其余字节就是具体的像素值和标签值。对做工程的人来说与其手写解析器不如直接用工具库我一般推荐两种方式一是用 PyTorch 的torchvision.datasets.MNIST一行代码解决下载和预处理二是用 TensorFlow/Keras 的数据集接口或者用 scikit-learn 提供的fetch_openml。# 方式1使用 scikit-learn 获取数据传统机器学习路线常用 from sklearn.datasets import fetch_openml # 下载 MNIST 数据集返回 DataFrame 格式 X, y fetch_openml(mnist_784, version1, return_X_yTrue, as_frameFalse, parserpandas) # 像素值转浮点并归一化到 [0,1]这是几乎所有模型的通用前置步骤 X X.astype(float32) / 255.0 # 将标签转成整数OpenML 返回的是字符串格式 y y.astype(int) print(f数据集形状: {X.shape}标签形状: {y.shape})这段代码的逻辑很直白fetch_openml会把 MNIST 自动下载到本地缓存目录as_frameFalse时返回 Numpy 数组而不是 DataFrame内存占用更小。parserpandas是 scikit-learn 新版绕开PyArrow警告的推荐写法如果你用的版本较老不传这个参数也能跑。归一化到 0-1 的原因是几乎所有梯度类模型都对量纲敏感灰度值 255 和 0 之间差太大会让权重更新出现“震荡”。第一次跑通时不需要做标准化但如果你想冲更高的准确率建议后面再对特征做标准化。2.2 一个不需要深度学习的基线KNN、逻辑回归和随机森林很多人听到“手写数字识别”就默认要用卷积神经网络这是一个误区。在无卷积的设定下传统机器学习方法依然有很强的竞争力。KNN 在原始像素特征上就能到 97% 左右逻辑回归约 92%随机森林大概能到 96% 到 97% 之间。这几个模型的共同点是训练快、调参少、可解释性强非常适合用来验证整个数据流水线是否正确。我们先从逻辑回归跑起它是最简单的“能跑通”的模型也是之后所有算法对比的基准线。这个基线的价值在于如果后续你的深度模型连逻辑回归都打不过那大概率不是模型问题是数据处理流程出了问题。from sklearn.linear_model import LogisticRegression from sklearn.model_selection import train_test_split from sklearn.metrics import accuracy_score import time # 数据太大时可以先抽样训练集取 5000 张跑基线 X_train, X_test, y_train, y_test train_test_split(X, y, train_size5000, test_size10000, random_state42) start time.time() # solverlbfgs 适合中小数据集max_iter 给足防止不收敛 clf LogisticRegression(solverlbfgs, max_iter1000, C0.1) clf.fit(X_train, y_train) end time.time() train_acc clf.score(X_train, y_train) test_acc accuracy_score(y_test, clf.predict(X_test)) print(f训练耗时: {end - start:.2f}s | 训练集准确率: {train_acc:.4f} | 测试集准确率: {test_acc:.4f})这里的参数说明值得多写几句。solverlbfgs是逻辑回归在非稀疏小数据上的默认选择收敛稳定C是正则化强度的倒数C 越小正则化越强MNIST 像素维度有 784 维样本只有 5000 张C 取 0.1 比默认的 1.0 更不容易过拟合。max_iter1000是因为像素特征没有归一化时迭代次数需求会变高虽然我们已经做了 /255.0 的归一化但标准的 Z-Score 标准化没有做给足迭代次数能减少“ConvergenceWarning”的打扰。逻辑回归只是热身真正能打的是随机森林。RF 对特征缩放不敏感能直接吃原始像素而且抗过拟合能力强。它的缺点是推理速度慢、模型体积大但在 28x28 的小图上完全不是问题。from sklearn.ensemble import RandomForestClassifier # n_estimators 控制树林规模这里 100 棵足够稳定n_jobs-1 用满所有CPU核 rf RandomForestClassifier(n_estimators100, max_depth12, n_jobs-1, random_state42) rf.fit(X_train, y_train) test_pred rf.predict(X_test) test_acc accuracy_score(y_test, test_pred) print(f随机森林测试集准确率: {test_acc:.4f}) # 查看特征重要性第几个像素对分类最重要 importance rf.feature_importances_.reshape(28, 28)随机森林的max_depth12是一个折中参数MNIST 的像素噪声很多树太深会记住噪声太浅又学不到位。特征重要性重新 reshape 成 28x28 后输出一个热力图能很直观看到模型主要在看图片中心区域边缘像素几乎不重要。这个可视化在写文档说明时非常好用能从“我跑了准确率”升级到“我理解了模型在看哪里”。传统机器学习路线里KNN 也是常被比较的对象但我个人并不建议把它作为主力模型因为 MNIST 有 60000 条训练数据KNN 的推理要对全部样本计算距离速度非常慢而且没有可解释性。用一次感受一下然后放弃是比较务实的做法。3. 从原始像素到有效特征PCA 降维与 HOG 特征的实际效果3.1 784 维像素直接进模型问题在哪里如果只用原始像素喂给逻辑回归或 SVM效果也不会太差但有三个问题让工程上不太能接受一是维度高导致训练慢、模型文件大二是像素之间相关性极强相邻像素的灰度值几乎一样信息冗余严重三是模型根本不知道“数字的形状”是什么只能靠像素值硬猜。特征工程的思路就是人工或半自动地从像素中提取“形状”信息。对 MNIST 这种简单数据集特征工程有三个常用方向PCA 降维、HOG 方向梯度直方图、图像缩放和重心对齐。其中重心对齐是最容易被忽略的一个玄学操作但却非常有效——把整张图的质心挪到画布中心相当于做了最朴素的空间归一化能直接提升 1 到 2 个百分点的准确率。import numpy as np def center_digit(img): 根据亮像素的质心把数字挪到图像中心 # 计算质心坐标 threshold 0.1 coords np.argwhere(img threshold) if len(coords) 0: return img cy, cx coords.mean(axis0) dy int(round(14 - cy)) dx int(round(14 - cx)) # 用 np.roll 做位移效率比逐像素复制高很多 shifted np.roll(img, dy, axis0) shifted np.roll(shifted, dx, axis1) return shifted # 对全部训练和测试数据做对齐这里用X的前2000条做演示 X_aligned np.array([center_digit(x.reshape(28, 28)).reshape(784) for x in X[:2000]])这段代码里有一个工程细节np.roll做的是循环位移数字从左边缘移出去会从右边缘再进来严格来说边缘 0 像素会“穿模”。由于 MNIST 背景基本都是 0实际影响很小但如果背景有噪声最好改用cv2.warpAffine或者scipy.ndimage.shift做边界填充。参数threshold 0.1是判断哪些像素属于“数字前景”的门限灰度归一化后背景噪声通常小于 0.1调大这个值会更抗噪但过低可能会把浅色数字的笔锋丢掉。3.2 PCA 降维的拐点主成分数量怎么定PCA 的作用是把 784 维像素压缩到一个低维子空间。这里的“低维”不是随便选一个数字而是要看“累积解释方差比”。我一般选让累计解释方差超过 0.9 的维度数量MNIST 在像素归一化后通常 100 到 150 个主成分就能达到这个阈值。如果只保留 30 维准确率会急剧下降因为很多数字的区分信息在容易被忽略的小方差方向上。from sklearn.decomposition import PCA # 直接对原始像素做 PCA先不标准化因为像素已经到 [0,1] 区间 pca PCA(n_components100, random_state42) X_pca_train pca.fit_transform(X[:5000]) # 在训练集上 fit X_pca_test pca.transform(X[5000:7000]) # 在测试集上只 transform # 查看解释方差占比 explained_ratio pca.explained_variance_ratio_ print(f前100个主成分累计方差占比: {explained_ratio.sum():.4f}) print(f单个最大主成分占比: {explained_ratio[0]:.4f})这段代码里最关键的是fit_transform和transform的区别PCA 的均值和特征向量只在训练集上学习测试集一律用训练好的参数转换这是防止数据泄漏的标准做法。很多人在这里踩坑——对全部数据一起fit_transform导致测试集的信息提前参与训练评估出来的准确率虚高 1% 到 2%上线后立刻原形毕露。解释方差比的意义在于告诉你第一个主成分通常占 20% 左右前 100 个主成分能保留 90% 以上的信息这部分“信息量”没有损失在视觉上但模型训练速度能快好几倍。还有一个在 MNIST 上很有效的方案是 PCA 白化whitening。把whitenTrue打开后主成分被缩放到单位方差相当于做了二次去相关。加了白化的 PCA 特征喂给逻辑回归准确率能提升 1 个百分点左右但代价是特征不再有直观的视觉解释。如果你写文档说明时想配图展示“降维后人脸/数字长什么样”就不要开白化特征向量还是能 reshape 成 28x28 看的。3.3 HOG 特征把形状变成梯度方向统计HOG方向梯度直方图是传统计算机视觉里最有代表性的特征之一它的核心逻辑是把图像切成小格子在每个格子里统计梯度方向形成一个柱状图。这个特征对光照变化不敏感而且特别适合描述“笔画”这种边缘结构。MNIST 的数字本质上就是不同方向笔画的组合HOG 天然匹配。from skimage.feature import hog def extract_hog(img_28x28): 提取单张 28x28 图像的 HOG 特征 features hog( img_28x28, orientations9, # 梯度方向分桶数 pixels_per_cell(4, 4), # 每个格子 4x4 像素 cells_per_block(2, 2), # 每个 block 2x2 格子 block_normL2-Hys, # 块归一化方式 visualizeFalse ) return features # 对前2000张图提取特征观察维度 sample_hog extract_hog(X[0].reshape(28, 28)) print(fHOG 特征维度: {sample_hog.shape}) X_hog np.array([extract_hog(x.reshape(28, 28)) for x in X[:2000]])参数说明是这篇文章的重点之一。orientations9意味着把 0 到 180 度的梯度方向分成 9 个区间每个区间一个桶pixels_per_cell(4, 4)是每个格子包含 4x4 像素28x28 的图会被分成 7x7 个格子cells_per_block(2, 2)是每个归一化块包含 2x2 个格子块与块之间有重叠。这三个参数决定了最终的特征维度压缩格子尺寸或者增加方向桶数都会使维度翻倍。HOG 特征喂给线性 SVM 是传统视觉的经典搭配准确率能到 98% 以上是所有传统机器学习方法里表现最好的之一。而且这个组合有一个额外优势SVM 在小样本上的泛化能力极强从 5000 张训练图中得到的模型准确率已经接近用全部 60000 张训练的效果。相比之下随机森林在小样本上会明显弱化。这就是为什么在数据不够的场景里“HOG SVM”至今还是很多工业项目的兜底方案。4. 模型对比与参数调优让准确率从 92% 跑到 98% 的关键调整4.1 一张表看懂各模型的“收益天花板”做完特征工程后需要系统性地对比一次模型效果。我一般固定训练集 6000 张、测试集 10000 张分别验证原始像素、PCA 特征、HOG 特征三种输入下各个模型的表现。这样能直观看到当前瓶颈在“特征”还是“模型”。模型原始像素PCA(100维)HOG特征逻辑回归91.8%92.5%94.0%KNN(k5)96.5%93.0%95.5%随机森林96.8%95.2%96.0%线性SVM92.5%93.5%95.8%RBF-SVM95.0%96.0%97.5%几个结论直接写在这里原始像素下 KNN 的“记忆式”分类很强因为它对特征缩放不敏感但换了 PCA 特征后 KNN 反而掉点原因是 PCA 白化与距离度量之间有冲突。随机森林加原始像素已经是性价比最高的“懒人组合”不需要特征工程就能到 96% 以上。真要冲 98%就得上 HOG 加 RBF-SVM这是传统机器学习路线的最优解。后面会解释 RBF-SVM 的两个关键参数。4.2 手写代码跑通对比实验的骨架为了避免每次实验都复制粘贴我会写一个极简的评估函数把“特征转换 模型训练 交叉验证”串起来。这样后续换特征或换模型只需要改一行。from sklearn.svm import SVC from sklearn.model_selection import cross_val_score def evaluate_model(model, X_data, y_data, cv3): 快速交叉验证评估模型性能 scores cross_val_score(model, X_data, y_data, cvcv, scoringaccuracy, n_jobs-1) print(f{model.__class__.__name__}: {scores.mean():.4f} (/- {scores.std():.4f})) return scores.mean() # 用 HOG 特征集评估线性核 SVM X_hog_6000 np.array([extract_hog(x.reshape(28, 28)) for x in X[:6000]]) y_6000 y[:6000] linear_svm SVC(kernellinear, C0.1) rbf_svm SVC(kernelrbf, C5.0, gammascale) print(--- 线性SVM on HOG ---) evaluate_model(linear_svm, X_hog_6000, y_6000) print(--- RBF-SVM on HOG ---) evaluate_model(rbf_svm, X_hog_6000, y_6000)这里有两处参数值得展开。SVC默认自带 one-vs-one 的多分类策略对 10 类数字会训练 45 个二分类器所以训练耗时比随机森林慢不少gammascale是让算法根据特征维度自动算 gamma在特征维度比较高时比固定值更靠谱。cross_val_score的cv3在 6000 张图上已经能反映泛化水平如果嫌慢可以降到 2但不要用分层划分缺失的默认策略MNIST 各类样本均衡不分层也可以。4.3 RBF-SVM 调参C 和 gamma 是一对冤家RBF-SVM 最核心的两个参数是C和gamma。gamma控制单个训练样本的影响半径gamma 越大边界越复杂容易过拟合gamma 越小决策边界越平滑容易欠拟合。C控制对错误分类的惩罚力度C 越大训练时越不允许出错边界越紧同样容易过拟合。它们在 MNIST 上典型的“好参数”窗口是C在 1 到 10 之间、gamma在 0.001 到 0.01 之间具体数值需要用小规模网格搜索锁定。from sklearn.model_selection import GridSearchCV # 小规模网格搜索暴力扫一组参数 param_grid { C: [0.1, 1.0, 5.0], gamma: [0.001, 0.005, 0.01], } grid GridSearchCV( SVC(kernelrbf), param_grid, cv3, scoringaccuracy, n_jobs-1, verbose1 ) # 用前3000条HOG特征跑网格搜索速度可控 grid.fit(X_hog_6000[:3000], y_6000[:3000]) print(f最佳参数: {grid.best_params_}) print(f最佳交叉验证准确率: {grid.best_score_:.4f})网格搜索的套路是先用小数据粗扫锁定大区间后再细扫。直接在 60000 张训练集上做 3x3 网格搜索会导致 RBF-SVM 训练 45 个二分类器乘 9 组参数时间会膨胀到不可接受。我用这个策略时通常会先在 3000 张图上找到备选区间再逐步放大训练集最后在全部数据上用最佳参数重新训练。这种做法比“一次扫描全部数据”更符合工程节奏也更容易定位到底是数据量不够还是参数不对。5. MNIST 训练中的常见翻车现场现象、原因与对症解决5.1 下载 MNIST 数据集报 404 或连接超时现象是 torchvision 或 sklearn 在运行中弹出 HTTP 错误提示 Unable to download 或 404 Not Found。很多新手以为是自己网络问题其实根本原因是 MNIST 原始站点yann.lecun.com有时候不稳定而 PyTorch 的torchvision.datasets.MNIST默认从该站下载一旦源站临时下线就会 404。scikit-learn 的fetch_openml走的是 OpenML 镜像通常没问题但要装 pandas 解析器。解决方法是换镜像源或本地缓存。一个稳妥做法是先用fetch_openml下载一遍并存成.npz文件以后所有实验都从本地.npz加载不再碰网络。另一个思路是找 GitHub 上常见的 MNIST 镜像仓库下载原始 IDX 文件放到torchvision的root目录下的MNIST/raw文件夹里。这里有一条“后悔药”任何时候都先落盘一份原始 IDX 文件这份文件就是整个项目的“数据保险单”。import numpy as np # 第一次下载成功后立即保存为 NPZ后续加载不走网络 from sklearn.datasets import fetch_openml X, y fetch_openml(mnist_784, version1, return_X_yTrue, as_frameFalse, parserpandas) np.savez_compressed(mnist_784.npz, XX, yy.astype(int)) # 之后重新加载 data np.load(mnist_784.npz) X, y data[X], data[y] print(f本地缓存加载完成: {X.shape})5.2 HOG 特征计算时间长到怀疑人生现象是代码运行了十几分钟还停在特征提取这一步让人以为是卡死了。原因是对 60000 张训练图逐张调用skimage.feature.hog时Python 的 for 循环开销巨大每张图大约要 30 到 50 毫秒60000 张就是 30 到 50 分钟。这在一次完整的实验流程里属于“太慢但不至于崩溃”但会浪费大量等待时间。解决方法是启用多进程。最简单的方式是multiprocessing.Pool把图像分批分发到多个进程提取结果再合并。注意skimage的 HOG 是 CPU 密集操作多进程加速比多线程更有效因为 Python 的 GIL 会限制多线程执行纯计算任务。另一个优化是降低 HOG 输入分辨率先缩放到 20x20 再提取特征速度能提升近一倍准确率只掉 0.2% 左右。5.3 训练集准确率很高、测试集准确率差一大截现象是训练集准确率 99.5%测试集只有 96%显然过拟合了。原因常见有两个一是随机森林或 RBF-SVM 的参数过于激进比如max_depthNone或者gamma0.1二是特征工程时用了“全数据集 fit”的 PCA导致测试集信息泄漏测试准确率虚高或者反常。解决思路是先简化模型复杂度再用“训练集和测试集准确率的差值”作为辅助指标排查泄漏。我的习惯是建立一条“铁律”PCA 的均值、标准化参数、特征缩放参数一律只在训练集上拟合然后应用到验证集和测试集。Pandas 和 Numpy 操作都非常容易在这条铁律上“翻车”所以最好把特征工程封装成一个类一个fit方法和一个transform方法确保不会跨数据集泄漏。随机森林的剪枝参数min_samples_leaf5也能显著抑制过拟合代价是训练集准确率降到 98% 左右但测试集能回升到 96.8% 附近整体更健康。5.4 torchvision 下载 MNIST 报 SSL 证书错误现象是在公司内网或部分 Linux 服务器上torchvision.datasets.MNIST抛出SSL: CERTIFICATE_VERIFY_FAILED而不是 404。原因是服务器缺少根证书或系统时间不对导致 HTTPS 握手失败。解决办法不是关 SSL 验证而是下载原始文件手动放置到root/MNIST/raw目录并设置downloadFalse。这样可以完全避开网络握手流程。如果手动放置也嫌麻烦另一个稳定方案是使用fetch_openml完成下载后再切成 PyTorch 的 Dataset 格式。OpenML 走的是自身镜像稳定性要好得多。实际上 PyTorch 生态里很多人已经习惯用 OpenML 数据配合自定义 Dataset 类这样既绕开了 404 和 SSL 的坑又不影响后续DataLoader使用。6. 从传统机器学习切到 PyTorch 卷积网络什么时候值得换、怎么复用数据6.1 传统方法能到 98%为什么还要碰卷积如果目标只是交作业或写入门笔记HOG 加 SVM 已经绰绰有余。但如果你的目标是“为后续计算机视觉项目打底”就必须碰 PyTorch 版本的卷积网络。卷积网络的优势在于不用人工设计特征网络自己从像素中学习边缘、纹理和部件这是和“机器学习算法”最本质的差别。MNIST 上卷积网络随便跑就能到 99.2% 以上超过传统 ML 的最优结果但这并非重点——重点是你可以在几天内验证一套完整的训练流程数据加载、模型定义、损失函数、优化器、训练循环、评估函数。这些套路在《机器学习》课本后期内容里几乎不会涉及却是实际做项目时天天用的东西。从传统特征切到卷积时之前做的 HOG 特征全部可以弃用数据归一化也简化成了“(像素 / 255.0) - 0.5”这一个人人都会写的操作。但有一个旧经验可以保留重心对齐。对 CNN 来说重心对齐依然能带来 0.1% 到 0.2% 的提升虽然幅度不大但它给了网络一个“输入分布更一致”的空间。我曾经做过对比实验对齐后模型收敛速度快了约 20%。6.2 PyTorch 训练 MNIST 的最小骨架从加载到评估这里给出一个可以直接替换到自己项目里的最小 PyTorch 骨架。数据加载用torchvision完成网络用一个两层卷积加全连接的极小结构保证在 CPU 上跑也能在 3 分钟内完成一个 epoch让新手能即时看到反馈。import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader, TensorDataset from sklearn.datasets import fetch_openml # 加载数据并转成 PyTorch Tensor X, y fetch_openml(mnist_784, version1, return_X_yTrue, as_frameFalse, parserpandas) X X.astype(float32).reshape(-1, 1, 28, 28) / 255.0 y y.astype(int64) # 划分训练和验证集这里用 5000 张演示 X_train, X_val, y_train, y_val X[:5000], X[5000:6000], y[:5000], y[5000:6000] train_loader DataLoader(TensorDataset(torch.tensor(X_train), torch.tensor(y_train)), batch_size128, shuffleTrue) val_loader DataLoader(TensorDataset(torch.tensor(X_val), torch.tensor(y_val)), batch_size128, shuffleFalse) # 极小卷积网络两层卷积 全局池化 全连接 class SmallCNN(nn.Module): def __init__(self): super().__init__() self.features nn.Sequential( nn.Conv2d(1, 16, kernel_size3, padding1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(16, 32, kernel_size3, padding1), nn.ReLU(), nn.MaxPool2d(2), ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(32 * 7 * 7, 64), nn.ReLU(), nn.Linear(64, 10), ) def forward(self, x): return self.classifier(self.features(x)) model SmallCNN() criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr1e-3) # 训练 3 个 epoch for epoch in range(3): model.train() running_loss 0.0 for images, labels in train_loader: optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() print(fEpoch {epoch1}: Loss {running_loss / len(train_loader):.4f}) # 验证 model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in val_loader: outputs model(images) _, predicted torch.max(outputs.data, 1) total labels.size(0) correct (predicted labels).sum().item() print(f验证集准确率: {100 * correct / total:.2f}%)这段代码值得逐行解释关键点。reshape(-1, 1, 28, 28)是 PyTorch 的通道顺序C x H x W不是 TensorFlow 的 H x W x C很多从 Keras 转过来的人第一课就是在维度顺序上翻车。nn.CrossEntropyLoss在 PyTorch 里自带 softmax所以模型最后一层不需要额外加 Softmax。optim.Adam的默认学习率1e-3在这个小网络上够用但如果你加大网络需要降到3e-4以下不然损失容易“震荡”。整个训练循环里with torch.no_grad()是验证时的标准写法能减少一半内存占用和计算量。6.3 对比传统方法和 CNN 的边界什么时候该回头用 SVM学会 PyTorch 之后也并不意味着所有数据都该上 CNN。在样本量小于 2000 张的小数据集上HOG 加 SVM 往往比 CNN 更稳定因为 CNN 的数据需求量更大小样本上很快就会过拟合。而在 28x28 这种低分辨率图像上CNN 相对传统方法的优势并不像在 ImageNet 那种高分辨率大图上那么明显。如果要做实时性或嵌入式部署传统 ML 模型的推理耗时要低一个数量级模型体积也更小没有 GPU 的服务器上跑 SVM 比跑 CNN 舒服得多。所以我给团队的建议是先建立传统 ML 基线再上 CNN二者互补而不是二选一。长期来看掌握这条“传统 ML 深度学习”双路线才能在职场上应对类型参差不齐的视觉任务。毕竟真实项目里“数据太少”“没有 GPU”“要上嵌入式设备”才是常态MNIST 只是让这些经验能在最小规模上完整预演一遍。6.4 一个亲测有效的进阶验证用 softmax 输出做置信度曲线训练完模型后不要只停在准确率数字上。画一张“置信度直方图”能帮你发现很多隐藏问题如果大量预测的 softmax 概率集中在 0.9 以上说明模型过度自信对模糊样本的区分力可能不足如果集中在 0.4 到 0.5 附近说明模型欠置信。在 MNIST 上一个健康模型的置信度分布应该呈现“两极化”趋势——正确样本接近 1.0错误样本在 0.2 到 0.7 之间徘徊这样后续做“拒绝预测”时才有操作空间。具体实现只需要在验证集推理时把 softmax 值存成数组用matplotlib画直方图。这个方法在普通源代码项目里不常见但对于“文档说明”部分的高质量呈现非常有帮助。我每次带新人都要求交两张图一张准确率曲线、一张置信度分布图这两张图比任何文字描述都更能说明模型状态。整条 MNIST 路线走完你会发现自己不只会调包而是真正理解了“机器学习应用流程”里每一个环节为什么存在。希望这个完整拆解能帮到你。本文还有配套的精品资源点击获取
返回列表