拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

用MNIST串起机器学习基础:从环境搭建到模型部署全流程

用MNIST串起机器学习基础:从环境搭建到模型部署全流程

带过不少刚开始学机器学习基础的新人,我发现一个很普遍的现象:大家不是缺资料,是缺一条能把知识点串起来的主线。公式背了一堆,框架也调得动,可一旦模型效果不好,就完全不知道从哪里下手排查。这篇文章想跟你聊的,就是这条主线本身。我打算用MNIST这个数据集把机器学习基础从头到尾过一遍,从环境搭建、数据处理、训练循环,到评估、部署选型,把每个环节里最容易让人卡住的地方摊开来讲。无论你是学生、转行的工程师,还是自己折腾过几个教程但还没系统串过一遍的爱好者,这篇文章应该都能给你省下不少弯路。

1. 拿MNIST当第一课:为什么是它而不是别的数据集

1.1 一份数据吃透四条主线

MNIST是手写数字识别数据集,60000张训练图片、10000张测试图片,每张是28x28的灰度图,内容就是0到9这十个数字。很多人觉得它太老、太简单,不愿意花时间在上面,其实这是误解。机器学习基础的核心不是模型有多新,而是你要在一份可控的数据上,把“数据准备、模型定义、训练优化、评估调参”这四条线全部走通。

我见过太多人一上来就上图像分类大模型或者NLP任务,结果数据清洗和处理占用了一半时间,反而没有精力理解模型训练的本质。MNIST最难得的地方在于它足够干净,你不需要做复杂的数据清洗,不需要处理缺值、异常文本,甚至不需要做数据增强,就能把注意力全部放在学习流程上。

还有一点很实际:MNIST迭代极快。在普通CPU上训练一个简单的全连接网络,一个epoch也只需要几秒到十几秒,这让你可以大胆地改学习率、改网络层数、看各种操作对结果的影响。这种“快速闭环”对建立直觉非常重要。如果你第一个项目就跑需要半小时以上的数据集,你是不太可能愿意反复做实验的。

1.2 搜索热词里藏着的三个信号

我去查了下最近的网络热词,发现“mnist for ml beginners”出现频率一直很高。这说明即使现在各种教程满天飞,MNIST依然是无数人入门时避不开的第一站。但光看搜索量其实看不到问题,真正值得关注的是新手搜索背后的困惑,根据我这些年看到的真实案例,大家通常卡在三个点上。

第一,把“跑通教程”当成了“学会”。Jupyter里把官方示例跑通,准确率显示99%,然后就没有然后了。代码不是自己写的,参数不知道为什么这么设,换一个数据集就彻底不会了。第二,版本错位导致的环境问题。网上的教程可能是两年前写的,Python版本、框架版本、CUDA版本都对不上,照着敲就是报错,于是大量时间浪费在排查环境上。第三,只看最终准确率,不看训练过程中的损失曲线。准确率到99%就以为万事大吉,完全不知道训练过程中发生过过拟合或者梯度爆炸。这篇文章后面会把这几个坑逐一展开。

2. 环境搭建的版本选择题:Python、CUDA与框架的三角关系

2.1 我推荐的入门组合和理由

机器学习基础阶段,环境选型只有一个原则:让版本问题尽量少来打扰你。我个人最推荐的新手组合是Python 3.10或者3.11、PyTorch 2.x的稳定版,如果你有NVIDIA显卡,再配上和显卡驱动匹配的CUDA版本。

为什么用PyTorch而不是TensorFlow?不是说TensorFlow不好,而是PyTorch的调试体验对新手更友好,报错信息相对直接,动态图模式下你可以用print随时看中间张量的形状,这对理解数据在模型里怎么流动特别有帮助。至于Python版本,不要图新鲜装最新的,PyTorch官方对Python 3.12、3.13的支持往往滞后,装完装不上依赖又得折腾半天。

这里给你一个我实际用过的版本组合参考:

组件推荐版本理由
Python3.10 / 3.11PyTorch官方wheel支持最稳定
PyTorch2.x稳定版动态图友好,社区资料最丰富
CUDA12.x(配合驱动)大部分新卡和编译版本都兼容
包管理pip / venv环境隔离,避免系统级污染

另外强烈建议你在项目目录下用虚拟环境,不要图省事直接全局装。我就是当年偷懒,把TensorFlow和PyTorch装在同一个全局环境里,结果OpenMP冲突,每次跑训练都会Crash,排查过程极其痛苦。现在凡是新项目,第一步永远是python -m venv venv。

2.2 硬件加速和“授权文件”话题为什么还不属于你

选环境的时候很多人会问:CPU能不能学机器学习基础?答案是完全可以。MNIST这种规模的数据,CPU训练一个全连接网络毫无压力,甚至卷积网络也能跑,就是慢一些。所以入门阶段不要为了GPU焦虑。

但有个热词我想专门提一下:很多人搜“Vivado ML 2023.1要不要授权文件”。这里涉及的是FPGA上的机器学习工具链。我理解大家的想法——既然机器学习最后要落地,不如一开始就往硬件加速上靠。这个思路本身没错,但顺序错了。机器学习基础阶段你连模型训练和评估的逻辑都没理顺,直接上FPGA工具链,等于还没学会开车就在研究发动机ECU调校。

Vivado ML这类FPGA开发工具确实涉及授权文件、版本兼容、芯片型号匹配一堆事情,这些属于部署和嵌入式优化阶段的问题。我的建议是:先把基础阶段的训练和评估跑扎实,等到真要部署到边缘设备时,再回头研究这些工具链。中间可以了解ONNX、量化、剪枝这些通用概念,它们比具体某个FPGA工具更接近机器学习基础的核心。

3. 把784个像素变成模型的输入:数据预处理里的小细节

3.1 归一化不是可选步骤

MNIST每张图是28x28的灰度图,展开成一个向量就是784个浮点数,这就是“784个像素”这个说法的由来。原始像素值范围是0到255,如果你直接把这个整数喂给模型,网络也能训练,但效果通常不理想。

原因在于神经网络里的权重初始化和梯度更新,都是围绕“数值在合理范围”这个假设设计的。当输入范围是0到255时,某些层的加权求和结果会变得非常大,梯度也容易被放大,训练就会不稳定。这就好比你要拧一个螺丝,工具明明是按毫米设计的,结果你拿了个米尺来拧,不是不能拧,是容易滑丝。

把像素值归一化到0到1范围内,是入门阶段性价比最高的操作。具体做法很简单,除以255就行。如果你想更讲究一点,可以算数据集的均值和标准差,做标准化。

from torchvision import transforms transform = transforms.Compose([ transforms.ToTensor(), # 转成Tensor,同时把0-255变成0-1 transforms.Normalize((0.1307,), (0.3081,)) # MNIST官方均值/标准差 ])

这里ToTensor已经帮你完成了除以255的操作,Normalize再做标准化。有些人嫌麻烦省略Normalize,其实也能跑,但加上之后收敛会更稳定。如果你以后处理更复杂的数据集,这个习惯能省很多事。

3.2 DataLoader里的shuffle和batch_size必须真正理解

数据准备好了,接下来就是怎么喂给模型。PyTorch里的DataLoader是绕不开的组件,但大多数新手只是照抄,不理解里面两个关键参数到底在干什么。

第一个是shuffle。训练时设shuffle=True,验证和测试时设shuffle=False。为什么要这样?因为训练时如果每个epoch都按固定顺序喂数据,模型可能会学到数据顺序里的虚假规律,哪怕这是不该学的。打乱顺序能让每个batch的分布更随机,梯度更新也更稳定。而验证时你需要的是稳定可复现的评估结果,所以不打乱。

第二个是batch_size。这个参数直接决定每次参数更新前你“看”多少张图。MNIST上我常用64或者128,太小比如1,梯度噪声大会抖得很厉害;太大比如1000,一个epoch梯度更新次数太少,收敛速度反而不理想。你可以直观理解为:batch_size是一个人一天看多少道题再总结一次规律,看太少总结的规律太碎,看太多总结的频率又太低。

from torch.utils.data import DataLoader from torchvision.datasets import MNIST train_dataset = MNIST(root="./data", train=True, download=True, transform=transform) train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True, num_workers=2) test_dataset = MNIST(root="./data", train=False, download=True, transform=transform) test_loader = DataLoader(test_dataset, batch_size=256, shuffle=False, num_workers=2)

num_workers这个参数我也多说一句,它代表用几个子进程来加载数据。Windows上设成0最省心,Linux和Mac可以设2或者4。不是越大越好,设太大容易把内存吃满,反而拖慢速度。这些参数以后在每个项目里都会遇到,早理解早受益。

4. 训练循环里最难解释的部分:loss不动、accuracy抖动怎么办

4.1 第一个训练循环按什么顺序写

不管用什么框架,训练循环的骨架是一样的。很多新手把代码抄下来能跑,但问一句“为什么这里要optimizer.zero_grad()”就答不上来。我先给你一个标准的PyTorch训练循环:

optimizer = torch.optim.Adam(model.parameters(), lr=1e-3) loss_fn = torch.nn.CrossEntropyLoss() for epoch in range(5): model.train() running_loss = 0.0 for images, labels in train_loader: optimizer.zero_grad() # 1. 清空上一次的梯度 outputs = model(images) # 2. 前向传播 loss = loss_fn(outputs, labels) # 3. 计算损失 loss.backward() # 4. 反向传播 optimizer.step() # 5. 更新参数 running_loss += loss.item() print(f"epoch {epoch+1}, loss: {running_loss/len(train_loader):.4f}")

这个循环里的每一步都有明确目的。尤其是optimizer.zero_grad(),很多人会漏掉它。PyTorch的梯度是累加的,如果不每次清空,上一轮的梯度会叠加到这一轮上,参数的更新方向就会错乱,训练基本不可能收敛。你可以把梯度想象成黑板上的算式,每解完一道题得先擦掉再写下一道,不清空就会写得密密麻麻看不清。

再强调一下model.train()这个状态调用。它会在训练模式下启用Dropout和BatchNorm的训练行为,验证时则要切换成model.eval()。新手最容易在这上面栽跟头:训练完直接拿模型做验证,没切eval,结果BatchNorm用的还是训练时的统计量,评估结果就莫名其妙变差。

4.2 损失曲线和准确率曲线到底怎么读

训练跑起来了,第二个经典困惑就是:为什么损失不降了?为什么准确率在抖?这时候不要慌,先看曲线形状再下结论。

正常情况下,训练初期损失会快速下降,然后逐渐变缓,准确率则在某个区间内小幅震荡。这个震荡是正常的,因为每个batch的数据分布不完全一样,梯度更新有噪声。MNIST上用Adam优化器、学习率1e-3,一般1到3个epoch损失就会明显下降,5个epoch准确率就能到98%以上。

我整理了一份常见现象的排查表,都是新手群里出现频率极高的问题:

现象可能原因优先排查方向
损失完全不动学习率过低或梯度没回传检查是否有backward(),调大学习率
损失变成NaN学习率过高或数值不稳定减小学习率,检查输入是否包含非法值
准确率长时间在低位抖动数据没归一化或shuffle有问题检查预处理,确认shuffle=True
训练损失下降但验证损失升高过拟合增加Dropout、减少网络层数或加数据增强
验证准确率忽高忽低batch_size太小或数据顺序影响调大batch_size,确认eval模式

还有个很常见的误区是“损失必须降到0”。不是的。交叉熵损失即使模型学得很好也不会是0,它反映的是预测分布和真实分布的差异。MNIST上损失降到0.05以下已经是很不错的水平,你真正应该关注的是验证集上能不能稳定复现高准确率,而不是死死盯着损失绝对值。

我自己有个习惯:训练时每跑完一个epoch都打印训练损失和验证准确率,而不是等全部跑完再看。这样一旦第二个epoch损失不降,我能立刻停下来调参,而不是白白浪费半小时。新手也建议养成这个习惯,比什么魔法参数都管用。

5. 准确率到99%之后:评估模型时比数字更重要的三件事

5.1 混淆矩阵比准确率诚实

MNIST是有名的“准确率虚高”数据集,随便一个简单网络都能跑出97%以上,这就导致很多人只看一个数字就觉得自己模型很牛。但准确率会把很多问题藏起来。比如模型对数字0识别得特别好,对8经常误判成3,整体准确率还是很高,因为你根本没细看每个类别的表现。

混淆矩阵是更诚实的评估工具。它是一个10x10的矩阵,行代表真实标签,列代表预测标签,对角线上的数字就是每个类别预测正确的数量。拿MNIST来说,7和9、3和8、4和9这些数字长得太像,混淆矩阵能一眼看出模型具体在哪些类别上犯糊涂。

用sklearn一行代码就能出图:

from sklearn.metrics import confusion_matrix y_true = all_labels y_pred = all_predictions cm = confusion_matrix(y_true, y_pred) print(cm)

打印出来你就能清晰地看到:比如第7行第9列有个数字比较大,说明模型经常把7误判成9。这时候你就可以去查训练数据里7和9的样本是不是有某种特征差异没被学到,这种排查思路只有在混淆矩阵的引导下才走得通。

5.2 过拟合的识别:验证集不是摆设

机器学习基础里最重要的概念之一就是过拟合,但很多人对它的理解只停留在“训练集好测试集差”这句话上。在MNIST上,过拟合其实不那么容易发生,因为你用的小网络容量有限。但只要你把网络加深、加宽,或者训练轮次拉长,过拟合立刻就会出现。

怎样在训练过程中尽早发现过拟合?关键是盯住验证损失。如果训练损失还在下降,但验证损失开始反弹上升,这就是过拟合的典型信号。注意,验证损失上升不代表模型能力变差,而是模型开始“背诵”训练集了。它记住了训练集里那些数字的细节,包括噪声,结果遇到没见过的验证集反而更不自信。

应对过拟合的手段,按推荐顺序排列:增加Dropout层、降低网络容量、加数据增强、早停。MNIST上最立竿见影的是Dropout。在全连接层之间插一个nn.Dropout(0.2),验证准确率往往能再涨一点,模型对噪声的鲁棒性也更强。

5.3 不能拿测试集反复调参

这个原则很多新手会忽视,甚至很多教程也不强调。数据集划分有讲究:训练集用来更新参数,验证集用来调超参数,测试集是最后真正检验模型泛化能力的。如果你反复拿测试集去试模型效果,再回头调参,那测试集的信息实际上已经被你用进了模型选择过程,最终得到的“准确率”是有水分的。

打个比方,考试前你偷看了模拟卷的答案,虽然你分数很高,但如果正式考题和模拟卷不一样,你的水平就现出原形了。测试集就是你最后那场正式的“高考”,模拟卷是用来调整复习策略的,两者不能混。

在MNIST上实践时,官方已经帮你分好了训练集和测试集。我的做法是再从训练集里切出一小块当验证集,比如用random_split切出5000张。这样我调参时只看验证集结果,最后用测试集做一次终极验证,得出的数字才真正有参考价值。

6. 模型离开训练环境之后:导出、再训练与部署选型的现实问题

6.1 模型文件里到底存了什么

训练结束后,你会得到一个模型对象,但你不可能永远在Jupyter里用模型。你需要把它保存成文件,以后加载推理,或者部署到别的环境。这里就涉及一个基础但总被忽略的问题:模型文件里存的到底是什么。

PyTorch里最推荐的保存方式是只存state_dict,也就是模型的权重和偏置。这个文件本质上是一个Python字典,key是每一层的名字,value是那层参数的张量。你不需要保存整个模型对象,因为模型结构本身是写在代码里的,加载时先重建一个同结构模型,再把权重填进去。

# 保存 torch.save(model.state_dict(), "mnist_cnn.pth") # 加载 model = SimpleCNN() model.load_state_dict(torch.load("mnist_cnn.pth", map_location="cpu")) model.eval()

这里有个新手常犯的错:加载完模型直接拿去预测,忘了调用model.eval()。前面说过,不调用eval,模型里的Dropout和BatchNorm还处于训练行为,推理结果会不稳定。这个错误极其隐蔽,因为结果不是报错,而是准确率凭白掉一截,你查半天不知道原因。

6.2 从CPU推理到硬件加速:授权话题背后的实际门槛

保存好模型之后,接下来就是推理部署。最直接的方案是继续用PyTorch在CPU上加载模型跑推理,MNIST单张预测只要几毫秒,完全够用。如果你想部署到Web服务,可以再导出成ONNX格式,用ONNX Runtime加速推理。这里我不展开代码,因为属于进阶话题,但你要知道机器学习的落地路线一般是:PyTorch训练 -> 导出ONNX -> 用ONNX Runtime或TensorRT在目标设备上推理。

至于更高级的硬件加速,比如FPGA,那就要回到热词里提到的Vivado ML授权文件问题。我的态度很明确:到了FPGA这一步,你面对的不再是机器学习问题,而是硬件工程问题。你需要理解逻辑综合、时序约束、Resource占用、AXI总线通信,还要处理商业工具链的版本和许可证。这些和机器学习基础的“基础”二字已经没有太大关系了。授权要不要买、是选WebPack免费版还是完整版,属于另一个领域的问题。

入门阶段,我的建议是把部署目标定在CPU或者GPU推理就好。等你把MNIST在普通环境下跑得滚瓜烂熟,再决定要不要往FPGA方向深入,那时候研究Vivado ML的授权和版本问题,才是真正的对症下药。

7. 复盘:三个浪费过我大量时间的认知误区

7.1 “基础”不代表简单,但更不代表过时

我以前也犯过这个错,觉得MNIST太简单,直接跳过去学ResNet、Transformer那些看起来很厉害的东西。结果呢?看懂了结构图,却理解不了梯度为何消失,理解不了为什么ResNet要加跳接。后来回头老老实实把MNIST上的过拟合、学习率、BatchNorm这些基础概念一点点调实验调明白,再看那些复杂模型,豁然开朗。

如果你正在学机器学习基础,请一定把MNIST当成你的“实验田”。在这块田地上,你种什么都能快速看到结果。调大学习率,损失爆炸;调小学习率,训练变慢;加Dropout,验证损失下降;去掉归一化,训练抖动。这些直观体验比任何公式都更能建立你的模型直觉。

7.2 报错信息是第一手的教材

写训练脚本的过程一定会遇到无数报错,shape不匹配、内存不足、维度对不上。很多人一看到红字就慌,直接复制报错去搜。但别急着搜,先自己读一遍报错信息。PyTorch的报错已经标注了哪一行出的问题、张量的形状是什么、期望的形状是什么。学会读这个信息,比任何教程都更能帮你理解框架的数据流。

我记得有一次报错是RuntimeError: expected scalar type Float but found Double,原因是数据集标签的dtype是int64,和模型输出的浮点类型对不上。这种错误一旦理解了,以后遇到同类的就知道是类型问题。基础阶段最忌讳的是一路复制粘贴、一路遇错百度,那样跑通一百个项目也建立不起独立解决问题的能力。

7.3 框架只是工具,不是学习目标

最后一点是我最想强调的。PyTorch、TensorFlow、Keras,这些都只是工具。你换一个框架,模型结构、训练循环、数据处理逻辑全都要重写一遍,但底层的原理是不变的。不要在入门阶段今天想学PyTorch明天又想学JAX,先盯着一套框架把机器学习基础打牢。

我带人入门时经常说:框架是笔,机器学习是写作。你用哪支笔不重要,重要的是你能写出什么内容。你用Python能实现反向传播吗?你能从零手撸一个线性回归吗?你能在没有框架帮助的情况下解释清楚损失函数和梯度下降的关系吗?这些才是基础中的基础。框架封装得再好,也替代不了你对模型训练本身的掌控力。

如今再回看自己学机器学习的这条路,最大的体会就是:基础阶段的慢,其实就是快。MNIST上每一次看似无聊的调参实验、每一回报错后的源码阅读,都在为后来的复杂项目铺路。如果你能把这篇文章里提到的环境、预处理、训练循环、过拟合识别这些环节都亲手过一遍,并能在MNIST上自己调出一个达到98%以上准确率的模型,那你机器学习基础的地基已经算是打扎实了。接下来无论往计算机视觉、NLP还是模型部署方向发展,都会比直接啃复杂模型要轻松得多。

返回列表