发布时间:2026/9/4 17:12:59
【深度学习入门】PyTorch 零基础实战:全连接网络实现 MNIST 手写数字识别 文章目录一、MNIST数据集简介二、完整代码实现三、核心模块讲解3.1 数据集与DataLoader3.2 设备自动选择3.3 全连接神经网络模型搭建3.4 训练函数train()完整训练五步3.5 测试函数test()模型评估3.6 损失函数与优化器3.7 PyTorch常用损失函数一览四、运行结果说明五、常见踩坑总结六、总结一、MNIST数据集简介MNIST手写数字数据集是深度学习入门最经典数据集一共70000张28×28单通道灰度图片训练集60000张用来训练神经网络权重测试集10000张用来评估模型泛化能力每张图片对应标签0‑9代表手写数字类别像素原始范围0‑255代码中转换为张量后归一化到0~1之间。二、完整代码实现 MNIST手写数字数据集介绍 一共70000张灰度图片60000张训练集10000张测试集。 图片大小28×28像素单通道灰度图。 fromtorchimportnn# 导入神经网络模块fromtorch.utils.dataimportDataLoader# 数据加载器把数据集分批次打包fromtorchvisionimportdatasets# torchvision内置数据集库MNIST在这里fromtorchvision.transformsimportToTensor# 转换器PIL图片 → PyTorch张量Tensor构建训练数据集对象training_datadatasets.MNIST(rootdata,# 数据集本地存放文件夹在当前项目目录下生成data文件夹trainTrue,# True代表读取训练集(6万张)downloadTrue,# True本地没有文件就联网下载已有文件直接跳过下载transformToTensor(),# 将图片转为Tensor张量像素0~255归一化到0.0~1.0)构建测试数据集对象test_datadatasets.MNIST(rootdata,# 和训练集存到同一个data目录trainFalse,# False读取测试集(1万张)用来评估模型效果downloadTrue,transformToTensor(),)importmatplotlib.pyplotasplt# 创建画布可视化查看9张手写数字图片figureplt.figure()foriinrange(9):img,labeltraining_data[i]# 取出第i张图片和对应的数字标签(0~9)figure.add_subplot(3,3,i1)# 创建3行3列子图依次摆放图片plt.title(label)# 子图标题显示真实数字标签plt.axis(off)# 关闭坐标轴只看图片plt.imshow(img.squeeze(),cmapgray)# squeeze去掉通道维度gray灰度图显示aimg.squeeze()plt.show()# 弹出图片窗口# DataLoader数据集分批次每一批64张图片train_dataloaderDataLoader(training_data,batch_size64)test_dataloaderDataLoader(test_data,batch_size64)# 打印一批数据的shape看懂数据维度forX,yintest_dataloader:print(fShape of X [N, C, H, W]:{X.shape})# N批次大小C通道H高W宽print(fShape of y:{y.shape}{y.dtype})# y标签的shape和数据类型break# 判断设备优先cuda(GPU)苹果设备mps都没有就用cpudevicecudaiftorch.cuda.is_available()elsempsiftorch.backends.mps.is_available()elsecpuprint(fUsing{device}device)# 自定义神经网络类继承nn.ModuleclassNeuralNetwork(nn.Module):def__init__(self):super().__init__()# 调用父类nn.Module的初始化self.flattennn.Flatten()# Flatten展平层把28*28图片拉直成一维向量784self.hidden1nn.Linear(28*28,128)# 第一层全连接输入784输出128神经元self.hidden2nn.Linear(128,256)# 第二层全连接输入128输出256神经元self.outnn.Linear(256,10)# 输出层输入256输出10个类别(数字0‑9)# 前向传播定义数据流动路线函数名forward固定defforward(self,x):xself.flatten(x)# 将图片展平xself.hidden1(x)# 第一层全连接计算xtorch.sigmoid(x)# sigmoid激活函数引入非线性xself.hidden2(x)# 第二层全连接计算xtorch.sigmoid(x)# sigmoid激活函数xself.out(x)# 输出层得到10个类别的预测分数returnx# 创建模型对象并迁移到GPU/CPU设备上modelNeuralNetwork().to(device)print(model)# 打印网络结构 训练函数 dataloader训练数据加载器 model神经网络模型 loss_fn损失函数 optimizer优化器 deftrain(dataloader,model,loss_fn,optimizer):model.train()# 设置模型为训练模式开启dropout等训练专属逻辑本网络没用dropout但规范写法保留batch_size_num1# 记录当前是第几个batchforX,yindataloader:# 将图片数据、标签都搬运到GPU/CPU设备X,yX.to(device),y.to(device)predmodel.forward(X)# 前向传播得到预测结果lossloss_fn(pred,y)# 计算预测值与真实标签之间的损失optimizer.zero_grad()# 梯度清零上一轮的梯度要清空避免累加loss.backward()# 反向传播自动求各个参数的梯度optimizer.step()# 根据梯度更新神经网络权重w、bloss_valueloss.item()# 把tensor类型loss取出普通python数值ifbatch_size_num%1000:# 每100个batch打印一次损失print(floss:{loss_value:7f}[number:{batch_size_num}])batch_size_num1# 测试函数预留后续写测试、计算准确率逻辑deftest(dataloader,model,loss_fn):sizelen(dataloader.dataset)# 获取测试集总样本数量num_batcheslen(dataloader)# 获取测试集batch打包总个数model.eval()# 设置模型为评估模式停止权重更新test_loss,correct0,0# 初始化测试损失、正确样本计数withtorch.no_grad():# 关闭梯度计算不做反向传播节省显存forX,yindataloader:# 遍历测试集每一个批次X,yX.to(device),y.to(device)# 数据、标签迁移到GPU/CPUpredmodel.forward(X)# 前向传播得到预测输出test_lossloss_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(fTest result: \n Accuracy:{(100*correct)}%, Avg loss:{test_loss})loss_fnnn.CrossEntropyLoss()# 交叉熵损失函数多用于多分类任务optimizertorch.optim.SGD(model.parameters(),lr0.01)# SGD随机梯度下降优化器学习率0.01epochs10# 设置训练总轮数完整遍历训练集10次fortinrange(epochs):print(fEpoch{t1}\n-------------------------------)# 打印当前是第几轮训练train(train_dataloader,model,loss_fn,optimizer)# 执行一轮训练更新网络权重print(Done!)# 全部轮次训练完成提示test(test_dataloader,model,loss_fn)# 在测试集上评估模型效果计算loss和准确率三、核心模块讲解3.1 数据集与DataLoaderdatasets.MNISTtorchvision内置数据集downloadTrue自动下载数据集到本地data文件夹。ToTensor()把图片转为张量像素值归一化0‑1。DataLoader对数据集做分批次(batch)支持打乱、多线程读取。本例batch_size64每次给模型喂64张图片。数据维度格式[N, C, H, W]Nbatch批次大小C通道数H图片高度W图片宽度。MNIST灰度图C1。3.2 设备自动选择devicecudaiftorch.cuda.is_available()elsempsiftorch.backends.mps.is_available()elsecpu自动优先使用NVIDIA GPU(cuda)苹果硅芯片MPS最后降级CPU。模型和张量必须.to(device)搬运到对应设备才能运算。3.3 全连接神经网络模型搭建继承nn.Module是PyTorch自定义网络标准写法。nn.Flatten()将28×28图片展平成784维一维向量。nn.Linear全连接层实现y x W b yxWbyxWb。forward()函数必须定义描述数据前向流动路径不要手动调用模型对象(X)会自动调用forward。torch.sigmoid()激活函数引入非线性没有激活函数多层网络等价于单层线性模型。网络结构Flatten(784) → Linear(784→128) → sigmoid → Linear(128→256) → sigmoid → Linear(256→10)输出10维向量代表数字0‑9各个类别的得分。3.4 训练函数train()完整训练五步深度学习训练循环五大步骤前向传播pred model(X)得到预测输出计算损失loss loss_fn(pred,y)对比预测与真实标签差距梯度清零optimizer.zero_grad()梯度会累加每轮必须清空反向传播求梯度loss.backward()自动计算所有权重梯度参数更新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 损失函数与优化器损失函数 CrossEntropyLoss多分类任务首选内部集成LogSoftmaxNLLLoss输出层不需要额外加softmax。优化器 SGD随机梯度下降lr0.01为学习率控制每一步权重更新幅度。Epoch一轮epoch代表完整遍历全部训练集一次本例设置10轮完整训练。3.7 PyTorch常用损失函数一览损失函数使用场景CrossEntropyLoss多分类BCEWithLogitsLoss二分类MSELoss回归任务NLLLoss配合LogSoftmax多分类SmoothL1Loss/HuberLoss回归抗异常值四、运行结果说明程序运行首先自动下载MNIST数据集到./data文件夹。弹出matplotlib窗口展示9张手写数字样本。控制台打印张量维度、使用设备、网络结构。训练过程每100个batch打印loss损失值正常loss会逐步下降。10轮epoch训练结束后执行test函数输出测试集准确率和平均损失。提示如果使用sigmoid激活的简单全连接网络MNIST准确率一般可以达到95%左右。想要更高准确率可以改用ReLU激活、CNN卷积网络。五、常见踩坑总结忘记把model、X、y搬运到deviceCPU/GPU张量混合报错。训练循环忘记optimizer.zero_grad()梯度累加loss不收敛。测试阶段忘记model.eval()和torch.no_grad()显存占用高、评估结果异常。CrossEntropyLoss使用时自己额外加Softmax层会导致效果变差。forward不要手动调用model.forward(X)规范写法是model(X)。六、总结到这里我们就完整跑通了使用 PyTorch 解决 MNIST 手写数字识别的全流程。虽然这是一个入门项目但它涵盖了深度学习开发最核心的几个环节数据流转从 datasets 加载到 DataLoader 分批我们掌握了处理图像数据的标准姿势特别是 [N, C, H, W] 这个维度的概念以后处理任何视觉任务都离不开它。模型构建通过继承 nn.Module我们搭起了一个包含展平层、全连接层和激活函数的基础网络。这一步让你理解了数据是如何在网络中一层层流动并发生变换的。训练闭环这是最重要的一环。前向传播算预测、计算 Loss、梯度清零、反向传播求导、优化器更新参数——这五步法是深度学习的肌肉记忆必须烂熟于心。避坑与规范我们在代码中实践了设备自动切换GPU/CPU、训练/评估模式切换train/eval以及关闭梯度计算no_grad这些都是写出健壮代码的关键细节。虽然 MNIST 数据集规模较小、全连接网络结构相对简单但它犹如深度学习领域的“Hello World”麻雀虽小五脏俱全。掌握了这套标准训练范式你就具备了迁移学习的能力了。

相关新闻

2026/9/4 17:12:59

流量趋势 + SEO检测工具:旧文要不要改标题的数据判断法

标签:SEO检测工具 流量趋势 标题优化 数据驱动 SEO 检测 不只静态扫描,还看 流量趋势——回答「这篇还要不要投入修改成本」。 1. 三种典型曲线 缓降:内容老化,更新比新写划算平盘:长尾稳定,微调元数据即…

2026/9/4 17:12:59

单片机毕业设计-基于 STM32 或 51 单片机与 ESP8266 的物联网环境安防监测装置设计 基于 STM32 或 51 单片机的 WiFi 无线环境监测报警系统设计与实现(023806)

博主介绍:✌️码农一枚 ,专注于大学生项目实战开发、讲解和毕业🚢文撰写修改等。全栈领域优质创作者,博客之星、掘金/华为云/阿里云/InfoQ等平台优质作者、专注于嵌入式单片机,Java、小程序技术领域和毕业项目实战 ✌️…

2026/9/4 17:12:59

浏览器插件版 SEO检测工具:已发旧文还能怎么补救

标签:SEO检测工具 浏览器插件 旧文优化 墨衍 新文可以发布前检; 旧文库存 才是流量基本盘。墨衍 SEO 检测浏览器插件 支持打开任意 URL 即诊。 1. 插件适合的场景 浏览自己的 历史高 PV 文,顺手看结构是否老化读竞品时 对比元数据写法内链修…

2026/9/4 18:13:08

基于SpringBoot的企业知识库问答系统毕业设计项目源码

温馨提示:本人主页置顶文章(点我)开头有 CSDN 平台官方提供的学长联系方式的名片! 温馨提示:本人主页置顶文章(点我)开头有 CSDN 平台官方提供的学长联系方式的名片! 温馨提示:本人主页置顶文章(点我)开头有 CSDN 平台…

2026/9/4 18:13:08

带宽:GPU 的粮道在哪里堵住了(发烫优化系列 · 第 2 篇)

上一篇我们说过一个反直觉的结论:手机发烫,烧的不是"算力",而是"搬运"——Arm 官方给出的数字是 DRAM 访问每 GB/s 要付 80–100mW 的电。 这一篇把这个结论的物理基础彻底讲透:带宽到底是什么、为什么在手机…

2026/9/4 18:13:08

基于SpringBoot的汽车咨询管理系统的设计与实现毕业设计项目源码

温馨提示:本人主页置顶文章(点我)开头有 CSDN 平台官方提供的学长联系方式的名片! 温馨提示:本人主页置顶文章(点我)开头有 CSDN 平台官方提供的学长联系方式的名片! 温馨提示:本人主页置顶文章(点我)开头有 CSDN 平台…

2026/9/4 18:13:08

K8S LVM扩容-【20260902】001篇

文章目录 一、根目录扩容(剩余约 95GiB 全部给 `/`) 1.1 新建分区 /dev/sda3(类型 LVM) 1.2 让内核重读分区表 1.3 扩容 LVM 与文件系统 二、部署 NFS Server 2.1 修复 yum 源(CentOS 7 EOL 必做) 2.2 创建共享目录 2.3 配置 /etc/exports 2.4 固定端口(便于防火墙放行)…

2026/9/4 18:13:08

全桥电路与PWM控制:从逆变到整流的统一硬件平台解析

/* 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 18:08:08

工业视觉实战:基于全连接神经网络的喷码字符识别系统构建

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

2026/9/3 18:28:26

vSound小提琴数字处理器实操指南:从接线到演出的完整配置

电小提琴或者原声小提琴插电演出,第一个绕不开的坎就是声音难听。原声琴的共鸣和空气感一旦进了拾音器,出来的往往是一坨干瘪、发尖、带着奇怪塑料味的信号。我当初第一次把琴接上乐队调音台,直接被主唱吐槽"你这声音像在锯钢丝"。…

2026/9/3 14:29:47

传感器接口IC如何攻克生物化学传感的微弱信号难题?

1. 从电极到比特流:为什么生物化学传感必须依赖专用接口IC 做生物化学传感的人都有过类似的经历:明明传感器本身性能很好,信号输出却一塌糊涂——噪声大、漂移明显、重复性差,怎么调都达不到预期。很多时候问题并不在传感器&#…

2026/9/3 14:30:35

STM32F411CEU6多通道ADC采集:扫描模式+DMA实现详解

1. 多通道 ADC 的用武之地把“Multichannel ADC”和“STM32F411CEU6”这两个关键字放在一起,其实就是嵌入式开发里最常遇到的一类需求:用一块不算贵的 MCU,同时采集多路模拟信号。STM32F411CEU6 是 48 引脚的 Cortex-M4F 主控,主频…

2026/9/4 0:00:58

STM32H743 SPI从机DMA双缓冲通信实战

简介:本资源是面向嵌入式开发工程师与STM32进阶学习者的SPI DMA双机通信从机端完整实现方案,聚焦STM32H743高性能Cortex-M7单片机在工业控制与高速数据交互场景下的从机通信开发痛点。压缩包含1355个文件,主体为599个C源码与321个头文件&…

2026/9/4 0:00:58

CPU开盖降温教程:20元成本让温度直降30度的原理与实践

最近很多朋友都在抱怨,自己的电脑一到夏天就变成"烤箱",玩游戏时CPU温度动不动就飙到90度以上,风扇噪音堪比直升机。更让人头疼的是,明明配置不错,却因为高温降频导致性能大打折扣。如果你也遇到了类似问题&…

2026/9/4 0:00:58

ArkTS 表单工程:场地预约页的三态场次 Grid 与校验

ArkTS 表单工程:场地预约页的三态场次 Grid 与校验 App 14「运动场地预约」场地 Tab(Func1Tab),是整 App 交互最丰富的页面——场地横向切换 三色图例 渐变预约预览卡 快捷模板 今日场次 Grid(可选/已选/已满三态&…

2026/9/3 20:43:36

USB Type-C PCB布局分区设计:电源、高速信号与PD协议全攻略

做硬件这行,Type-C接口算是典型的“看着简单,做起来全坑”的东西。光引脚就24个,高低速信号、电源、控制线全部塞在一个小小的连接器里,如果PCB布局不做规划,打样回来基本就是“插上没反应”、“高速掉线”、“静电一打…

2026/9/3 17:51:43

系统编程学习原型如何补齐稳定性边界

系统编程学习原型如何补齐稳定性边界预算有限时&#xff0c;我先优化明显多余的复制&#xff0c;而不是猜测性地换容器。用借用传递只读数据通常就能减少分配&#xff1a; fn parse(line: &str) -> Result<Item, Error> { /* ... */ }用基准确认热点确实在分配&am…

2026/9/3 21:06:57

雨花区哪家财务公司代理记账比较好?

在雨花区&#xff0c;企业处理财税事务常常面临诸多挑战&#xff0c;选择一家靠谱的财务公司至关重要。湖南巨勤财务管理咨询有限公司就是本地正规实体财税服务机构&#xff0c;深耕本地工商财税行业多年&#xff0c;熟悉当地工商局、税务局最新政策与申报流程。主营公司注册、…