news 2026/9/4 17:12:54

【深度学习入门】PyTorch 零基础实战:全连接网络实现 MNIST 手写数字识别

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
【深度学习入门】PyTorch 零基础实战:全连接网络实现 MNIST 手写数字识别

文章目录

  • 一、MNIST数据集简介
  • 二、完整代码实现
  • 三、核心模块讲解
    • 3.1 数据集与DataLoader
    • 3.2 设备自动选择
    • 3.3 全连接神经网络模型搭建
    • 3.4 训练函数train():完整训练五步
    • 3.5 测试函数test():模型评估
    • 3.6 损失函数与优化器
    • 3.7 PyTorch常用损失函数一览
  • 四、运行结果说明
  • 五、常见踩坑总结
  • 六、总结

一、MNIST数据集简介

MNIST手写数字数据集是深度学习入门最经典数据集:

  1. 一共70000张28×28单通道灰度图片
  2. 训练集:60000张,用来训练神经网络权重
  3. 测试集:10000张,用来评估模型泛化能力
  4. 每张图片对应标签0‑9,代表手写数字类别
  5. 像素原始范围0‑255,代码中转换为张量后归一化到0~1之间。

二、完整代码实现

''' MNIST手写数字数据集介绍: 一共70000张灰度图片:60000张训练集,10000张测试集。 图片大小:28×28像素,单通道灰度图。 '''fromtorchimportnn# 导入神经网络模块fromtorch.utils.dataimportDataLoader# 数据加载器,把数据集分批次打包fromtorchvisionimportdatasets# torchvision内置数据集库,MNIST在这里fromtorchvision.transformsimportToTensor# 转换器:PIL图片 → PyTorch张量Tensor'''构建训练数据集对象'''training_data=datasets.MNIST(root="data",# 数据集本地存放文件夹,在当前项目目录下生成data文件夹train=True,# True代表读取训练集(6万张)download=True,# True:本地没有文件就联网下载;已有文件直接跳过下载transform=ToTensor(),# 将图片转为Tensor张量,像素0~255归一化到0.0~1.0)'''构建测试数据集对象'''test_data=datasets.MNIST(root="data",# 和训练集存到同一个data目录train=False,# False读取测试集(1万张),用来评估模型效果download=True,transform=ToTensor(),)importmatplotlib.pyplotasplt# 创建画布,可视化查看9张手写数字图片figure=plt.figure()foriinrange(9):img,label=training_data[i]# 取出第i张图片和对应的数字标签(0~9)figure.add_subplot(3,3,i+1)# 创建3行3列子图,依次摆放图片plt.title(label)# 子图标题显示真实数字标签plt.axis('off')# 关闭坐标轴,只看图片plt.imshow(img.squeeze(),cmap='gray')# squeeze去掉通道维度,gray灰度图显示a=img.squeeze()plt.show()# 弹出图片窗口# DataLoader:数据集分批次,每一批64张图片train_dataloader=DataLoader(training_data,batch_size=64)test_dataloader=DataLoader(test_data,batch_size=64)# 打印一批数据的shape,看懂数据维度forX,yintest_dataloader:print(f'Shape of X [N, C, H, W]:{X.shape}')# N批次大小,C通道,H高,W宽print(f'Shape of y:{y.shape}{y.dtype}')# y标签的shape和数据类型break# 判断设备:优先cuda(GPU),苹果设备mps,都没有就用cpudevice='cuda'iftorch.cuda.is_available()else'mps'iftorch.backends.mps.is_available()else'cpu'print(f'Using{device}device')# 自定义神经网络类,继承nn.ModuleclassNeuralNetwork(nn.Module):def__init__(self):super().__init__()# 调用父类nn.Module的初始化self.flatten=nn.Flatten()# Flatten展平层:把28*28图片拉直成一维向量784self.hidden1=nn.Linear(28*28,128)# 第一层全连接:输入784,输出128神经元self.hidden2=nn.Linear(128,256)# 第二层全连接:输入128,输出256神经元self.out=nn.Linear(256,10)# 输出层:输入256,输出10个类别(数字0‑9)# 前向传播,定义数据流动路线,函数名forward固定defforward(self,x):x=self.flatten(x)# 将图片展平x=self.hidden1(x)# 第一层全连接计算x=torch.sigmoid(x)# sigmoid激活函数,引入非线性x=self.hidden2(x)# 第二层全连接计算x=torch.sigmoid(x)# sigmoid激活函数x=self.out(x)# 输出层,得到10个类别的预测分数returnx# 创建模型对象,并迁移到GPU/CPU设备上model=NeuralNetwork().to(device)print(model)# 打印网络结构""" 训练函数 dataloader:训练数据加载器 model:神经网络模型 loss_fn:损失函数 optimizer:优化器 """deftrain(dataloader,model,loss_fn,optimizer):model.train()# 设置模型为训练模式,开启dropout等训练专属逻辑(本网络没用dropout,但规范写法保留)batch_size_num=1# 记录当前是第几个batchforX,yindataloader:# 将图片数据、标签都搬运到GPU/CPU设备X,y=X.to(device),y.to(device)pred=model.forward(X)# 前向传播,得到预测结果loss=loss_fn(pred,y)# 计算预测值与真实标签之间的损失optimizer.zero_grad()# 梯度清零:上一轮的梯度要清空,避免累加loss.backward()# 反向传播,自动求各个参数的梯度optimizer.step()# 根据梯度,更新神经网络权重w、bloss_value=loss.item()# 把tensor类型loss取出普通python数值ifbatch_size_num%100==0:# 每100个batch打印一次损失print(f'loss:{loss_value:>7f}[number:{batch_size_num}]')batch_size_num+=1# 测试函数,预留,后续写测试、计算准确率逻辑deftest(dataloader,model,loss_fn):size=len(dataloader.dataset)# 获取测试集总样本数量num_batches=len(dataloader)# 获取测试集batch打包总个数model.eval()# 设置模型为评估模式,停止权重更新test_loss,correct=0,0# 初始化测试损失、正确样本计数withtorch.no_grad():# 关闭梯度计算,不做反向传播,节省显存forX,yindataloader:# 遍历测试集每一个批次X,y=X.to(device),y.to(device)# 数据、标签迁移到GPU/CPUpred=model.forward(X)# 前向传播得到预测输出test_loss+=loss_fn(pred,y).item()# 累加本批次损失correct+=(pred.argmax(1)==y).type(torch.float).sum().item()# 统计本批次预测正确样本数a=(pred.argmax(1)==y)# 布尔张量:预测是否等于真实标签b=(pred.argmax(1)==y).type(torch.float)# 将布尔值转为float(1.0/0.0)test_loss/=num_batches# 计算测试集平均损失correct/=size# 计算测试集整体准确率print(f"Test result: \n Accuracy:{(100*correct)}%, Avg loss:{test_loss}")loss_fn=nn.CrossEntropyLoss()# 交叉熵损失函数,多用于多分类任务optimizer=torch.optim.SGD(model.parameters(),lr=0.01)# SGD随机梯度下降优化器,学习率0.01epochs=10# 设置训练总轮数,完整遍历训练集10次fortinrange(epochs):print(f"Epoch{t+1}\n-------------------------------")# 打印当前是第几轮训练train(train_dataloader,model,loss_fn,optimizer)# 执行一轮训练,更新网络权重print("Done!")# 全部轮次训练完成提示test(test_dataloader,model,loss_fn)# 在测试集上评估模型效果,计算loss和准确率

三、核心模块讲解

3.1 数据集与DataLoader

  • datasets.MNIST:torchvision内置数据集,download=True自动下载数据集到本地data文件夹。
  • ToTensor():把图片转为张量,像素值归一化0‑1。
  • DataLoader:对数据集做分批次(batch),支持打乱、多线程读取。本例batch_size=64,每次给模型喂64张图片。

数据维度格式[N, C, H, W]

  • N:batch批次大小;C:通道数;H:图片高度;W:图片宽度。MNIST灰度图C=1。

3.2 设备自动选择

device='cuda'iftorch.cuda.is_available()else'mps'iftorch.backends.mps.is_available()else'cpu'

自动优先使用NVIDIA GPU(cuda),苹果硅芯片MPS,最后降级CPU。模型和张量必须.to(device)搬运到对应设备才能运算。

3.3 全连接神经网络模型搭建

继承nn.Module是PyTorch自定义网络标准写法。

  1. nn.Flatten():将28×28图片展平成784维一维向量。
  2. nn.Linear:全连接层,实现y = x W + b y=xW+by=xW+b
  3. forward()函数:必须定义,描述数据前向流动路径,不要手动调用,模型对象(X)会自动调用forward。
  4. torch.sigmoid()激活函数,引入非线性;没有激活函数多层网络等价于单层线性模型。

网络结构:

Flatten(784) → Linear(784→128) → sigmoid → Linear(128→256) → sigmoid → Linear(256→10)

输出10维向量,代表数字0‑9各个类别的得分。

3.4 训练函数train():完整训练五步

深度学习训练循环五大步骤:

  1. 前向传播pred = model(X)得到预测输出
  2. 计算损失loss = loss_fn(pred,y)对比预测与真实标签差距
  3. 梯度清零optimizer.zero_grad(),梯度会累加,每轮必须清空
  4. 反向传播求梯度loss.backward()自动计算所有权重梯度
  5. 参数更新optimizer.step()使用梯度更新w、b权重

model.train():训练模式,部分层(Dropout、BN)会启用训练逻辑。

3.5 测试函数test():模型评估

  • model.eval():评估模式,关闭dropout、batchnorm训练行为。
  • with torch.no_grad():关闭梯度计算,节省内存,测试阶段不需要反向传播。
  • pred.argmax(1):取10维输出分数最大的下标,即为预测数字类别。
  • 统计correct预测正确样本数量,除以总样本得到准确率。

3.6 损失函数与优化器

  1. 损失函数 CrossEntropyLoss:多分类任务首选,内部集成LogSoftmax+NLLLoss,输出层不需要额外加softmax。
  2. 优化器 SGD随机梯度下降:lr=0.01为学习率,控制每一步权重更新幅度。
  3. Epoch:一轮epoch代表完整遍历全部训练集一次;本例设置10轮完整训练。

3.7 PyTorch常用损失函数一览

损失函数使用场景
CrossEntropyLoss多分类
BCEWithLogitsLoss二分类
MSELoss回归任务
NLLLoss配合LogSoftmax多分类
SmoothL1Loss/HuberLoss回归,抗异常值

四、运行结果说明

  1. 程序运行首先自动下载MNIST数据集到./data文件夹。
  2. 弹出matplotlib窗口展示9张手写数字样本。
  3. 控制台打印张量维度、使用设备、网络结构。
  4. 训练过程每100个batch打印loss损失值,正常loss会逐步下降。
  5. 10轮epoch训练结束后执行test函数,输出测试集准确率和平均损失。

提示:如果使用sigmoid激活的简单全连接网络,MNIST准确率一般可以达到95%左右。想要更高准确率可以改用ReLU激活、CNN卷积网络。

五、常见踩坑总结

  1. 忘记把model、X、y搬运到device,CPU/GPU张量混合报错。
  2. 训练循环忘记optimizer.zero_grad(),梯度累加loss不收敛。
  3. 测试阶段忘记model.eval()torch.no_grad(),显存占用高、评估结果异常。
  4. CrossEntropyLoss使用时,自己额外加Softmax层,会导致效果变差。
  5. forward不要手动调用model.forward(X),规范写法是model(X)

六、总结

到这里,我们就完整跑通了使用 PyTorch 解决 MNIST 手写数字识别的全流程。虽然这是一个入门项目,但它涵盖了深度学习开发最核心的几个环节:

  1. 数据流转:从 datasets 加载到 DataLoader 分批,我们掌握了处理图像数据的标准姿势,特别是 [N, C, H, W] 这个维度的概念,以后处理任何视觉任务都离不开它。
  2. 模型构建:通过继承 nn.Module,我们搭起了一个包含展平层、全连接层和激活函数的基础网络。这一步让你理解了数据是如何在网络中一层层流动并发生变换的。
  3. 训练闭环:这是最重要的一环。前向传播算预测、计算 Loss、梯度清零、反向传播求导、优化器更新参数——这"五步法"是深度学习的肌肉记忆,必须烂熟于心。
  4. 避坑与规范:我们在代码中实践了设备自动切换(GPU/CPU)、训练/评估模式切换(train/eval)以及关闭梯度计算(no_grad),这些都是写出健壮代码的关键细节。

虽然 MNIST 数据集规模较小、全连接网络结构相对简单,但它犹如深度学习领域的“Hello World”,麻雀虽小五脏俱全。掌握了这套标准训练范式,你就具备了迁移学习的能力了。


版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/4 17:06:43

如果组长要求你主导项目中的分库分表,大致的实施流程是什么

考点分析:这道题表面在问"流程",实际在考察你是否具备从业务痛点出发、完成架构设计与落地的全局能力。面试官重点会看以下几点: 能否给出从评估、设计、开发、迁移到上线运维的完整实施链路,而不是只背名词&#xff1b…

作者头像 李华
网站建设 2026/9/4 17:05:12

三极管实战指南:从原理到驱动电路,轻松驾驭半导体开关

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/4 17:01:48

一个适配器处理多任务:持续学习中的任务条件特征变换

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/4 17:00:12

免签支付系统实战:从通知监听到安全部署的完整架构解析

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/4 16:59:38

秋叶ComfyUI V17中文整合包:一键部署与AI绘图入门指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/4 16:59:13

AI模型一键整合包本地部署指南:从环境准备到功能测试

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华