🔄Update: [Lab-2] 数据装载

This commit is contained in:
2025-11-02 19:37:12 +08:00
parent 1e1f9644fe
commit d9d2293999
2 changed files with 67 additions and 0 deletions
View File
+67
View File
@@ -0,0 +1,67 @@
import torch
import torchvision
import matplotlib.pyplot as plt
from torchvision import transforms
from torch.utils.data import random_split, DataLoader
def load_and_preview_mnist(data_size=100, batch_size=16, seed=42, root='../data'):
"""
加载并预览 MNIST 数据集。
Parameters:
- data_size (int): 加载的数据条目数。
- batch_size (int): 每批次处理的数据数量。
- seed (int): 随机种子,用于保证数据切割的可重复性。
- root (str): 数据集保存的路径。
"""
# 数据预处理
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,))
])
# 载入数据集
train_dataset = torchvision.datasets.MNIST(
root=root,
train=True,
download=True,
transform=transform
)
# 截取指定数量的数据
data_train, _ = random_split(
dataset=train_dataset,
lengths=[data_size, len(train_dataset) - data_size],
generator=torch.Generator().manual_seed(seed)
)
# 加载数据集
train_loader = DataLoader(
dataset=data_train,
batch_size=batch_size,
shuffle=True
)
# 数据预览
images, labels = next(iter(train_loader))
# 将图像拼接为网格
img_grid = torchvision.utils.make_grid(images, nrow=batch_size // 4)
img = img_grid.numpy().transpose(1, 2, 0)
std = [0.5]
mean = [0.5]
img = img * std + mean
# 显示图像
print([labels[i] for i in range(batch_size)])
plt.imshow(img, cmap='gray')
plt.axis('off')
plt.show()
if __name__ == '__main__':
# 使用默认参数调用函数
load_and_preview_mnist()
# 或者使用自定义参数
# load_and_preview_mnist(data_size=200, batch_size=8, seed=123, root='./data')