目录
简介
一、PyTorch框架认识
1. Tensor张量
2.自动求导(Autograd)
3.构建神经网络(nn模块)
二、MNIST 手写数字识别
1.导入必要的库
2. 加载 MNIST 数据集
3. 数据可视化
4. 数据加载器
5. 设备配置
6. 构建神经网络模型
7. 训练和测试函数
8.多轮模型训练
三、模型优化
1. 损失函数(Loss Function)
2. 优化器(Optimizer)
简介
欢迎来到深度学习系列课程的第三课!本节课将聚焦PyTorch 框架在经典任务 ——MNIST 手写数字识别中的实践应用,带您从零搭建一个完整的神经网络模型,深入理解深度学习模型从构建、训练到评估的全流程。
MNIST 数据集作为深度学习入门的 “Hello World”,包含 70000 张 28×28 像素的手写数字灰度图像(60000 张训练集 + 10000 张测试集),任务目标是准确识别图像中的数字(0-9)。本节课将以该数据集为载体,手把手教您运用 PyTorch 实现核心步骤:首先,学习如何使用 PyTorch 的torchvision库加载并预处理 MNIST 数据,包括数据标准化、批量加载等关键操作,为模型训练做好数据准备;其次,详细讲解神经网络的搭建逻辑,从定义包含输入层、隐藏层、输出层的全连接网络结构,到选择合适的激活函数(如 ReLU)和损失函数(交叉熵损失),让您理解每一层的作用与参数设计思路;接着,深入模型训练流程,涵盖优化器(如 SGD、Adam)的配置、训练循环的编写(前向传播计算预测值、反向传播更新参数),以及如何监控训练过程中的损失变化与准确率提升;最后,学习使用测试集评估模型性能,分析模型在 unseen 数据上的泛化能力,并通过实际案例演示模型的预测过程,直观感受手写数字识别的效果。
一、PyTorch框架认识
1. Tensor张量
在PyTorch中,张量(Tensor)是核心数据结构,它是一个多维数组,用于存储和变换数据。张量类似于Numpy中的数组,但具有更丰富的功能和灵活性,特别是在支持GPU加速方面。
定义与特性 多维数组:张量可以看作是一个n维数组,其中n可以是任意正整数。它可以是标量(零维数组)、向量(一维数组)、矩阵(二维数组)或具有更高维度的数组。 数据类型统一:张量中的元素具有相同的数据类型,这有助于在GPU上进行高效的并行计算。 支持GPU加速:PyTorch中的张量可以存储在CPU或GPU上,通过将张量转移到GPU上,可以利用GPU的强大计算能力来加速深度学习模型的训练和推理过程。 创建方式 直接使用torch.tensor():根据提供的Python列表或Numpy数组创建张量。 下载数据集时:transform=ToTensor()直接将数据转化为Tensor张量类型。
2.自动求导(Autograd)
# 定义一个 tensor,并设置 requires_grad=True
x = torch.ones(2, 2, requires_grad=True)
print(x)
# 定义一个简单运算
y = x + 2
z = y * y * 3
out = z.mean()
# 反向传播计算梯度
out.backward()
print(x.grad)
- 注意:计算图在反向传播后默认会释放,如果需要多次反向传播,需要设置 retain_graph=True。
3.构建神经网络(nn模块)
nn.Module:所有神经网络模型都需要继承该类。(下面有具体的构建方法)
- 层级组合:可以将多层组合在一起,形成更复杂的网络结构。
二、MNIST 手写数字识别

对于mnist数据集的数据照片如下图,我将详细说明每一步的内容

1.导入必要的库
import torch
from torch import nn
from torch.utils.data import DataLoader
from torchvision import datasets
from torchvision.transforms import ToTensor
- torch:PyTorch 的主库,提供了张量操作和深度学习的基本功能
- nn:神经网络模块,包含了各种层和损失函数
- DataLoader:用于数据加载和批处理的工具
- datasets:包含了常用的数据集,这里使用 MNIST
- ToTensor:将图片转换为 PyTorch 张量的转换工具
2. 加载 MNIST 数据集
# 下载训练集
train_data = datasets.MNIST(
root=\’data\’,#数据集的根目录,下载位置
train=True, #如果为True,则从training.pt创建数据集,否则从test.pt创建数据集
download=True,# 为





