MNIST 代码解析 (4)下载、加载数据
2024-03-21 15:21:28
# 4 下载,加载数据
from torch.utils.data import DataLoader
# 下载数据集
train_set = datasets.MNIST("data", train=True, download=True, transform=pipeline)
test_set = datasets.MNIST("data", train=False, download=True, transform=pipeline)
# 加载数据
train_loader = DataLoader(train_set, batch_size=BATCH_SIZE, shuffle=True) # shuffle:打乱图片顺序
test_loader = DataLoader(test_set, batch_size=BATCH_SIZE, shuffle=True)
这段代码使用了PyTorch的数据加载和处理功能,特别是针对MNIST手写数字数据集。
MNIST是一个常用的入门数据集,包含了大量的手写数字图像(0到9),用于训练和测试机器学习模型。
代码解析
导入必要的库
form torch.utils.data import DataLoader
DataLoader是PyTorch中用于加载数据集的类,它提供了一个迭代器,可以按批次(batch)访问数据,支持自动批处理、采样、打乱数据和多线程数据加载等功能
2. 下载数据集
train_set = dataset.MNIST("data", train=True, download=True, transform=pipeline)
test_set = dataset.MNIST("data", train=True, download=True, transform=pipelien)
这里使用dataset.MNIST来下载MNIST数据集。
"data"是数据集的保存目录。
train=True表示下载的是训练集,train=False表示下载的是测试集。
download=True允许自动下载数据集。
transform=pipeline应用之前定义的预处理流程(pipeline)到每个图像上
3. 加载数据
train_loader = DataLoader(train_set, batch_size=BATCH_SIZE, shuffle=True)
test_loader = DataLoader(test_set, batch_size=BATCH-SIZE, shuffle=True)
使用DataLoader来加载数据集,其中batch_size=BATCH_SIZE定义了每个批次的大小,即每次迭代返回的图像数量。 什么是迭代?什么是迭代返回的图像数量?为什么要迭代返回数量?
shuffle=True表示在每个epoch开始时,数据将被打乱,这有助于模型学习使得泛化能力。 什么是模型泛化能力?
2024-03-21 15:36:44 十五分钟就写好了,真快
问题
什么是迭代?什么是迭代返回的图像数量?为什么要迭代返回数量?
什么是模型泛化能力?

浙公网安备 33010602011771号