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 十五分钟就写好了,真快

问题

什么是迭代?什么是迭代返回的图像数量?为什么要迭代返回数量?
什么是模型泛化能力?

posted @ 2024-03-21 15:36  MoonSheep|  阅读(131)  评论(0)    收藏  举报