From 1b58da9a47accb40ce7b9414a6758d7c3fa1e2e7 Mon Sep 17 00:00:00 2001 From: wonder Date: Sun, 2 Nov 2025 19:44:48 +0800 Subject: [PATCH] =?UTF-8?q?=20=F0=9F=94=84Update:=20[Lab-2]=20=E6=95=B0?= =?UTF-8?q?=E6=8D=AE=E5=A4=84=E7=90=86=E5=8F=98=E6=8D=A2?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Lab-2/src/data_transform.py | 88 +++++++++++++++++++++++++++++++++++++ 1 file changed, 88 insertions(+) create mode 100644 Lab-2/src/data_transform.py diff --git a/Lab-2/src/data_transform.py b/Lab-2/src/data_transform.py new file mode 100644 index 0000000..8a7eeb7 --- /dev/null +++ b/Lab-2/src/data_transform.py @@ -0,0 +1,88 @@ +import torch +import torchvision +import torchvision.transforms as transforms +import matplotlib.pyplot as plt +from torch.utils.data import DataLoader, random_split + + +def get_transform(): + """ + 返回 torchvision.transforms.Compose 对象,包含数据预处理变换操作。 + + Returns: + - transform (torchvision.transforms.Compose): 包含数据预处理变换的 Compose 对象。 + """ + return transforms.Compose([ + transforms.ToTensor(), # 将输入数据类型转换为 Tensor + transforms.Normalize((0.5,), (0.5,)) # 对数据进行标准化处理,均值和标准差均为 0.5 + ]) + + +def load_mnist(data_size=100, batch_size=16, seed=42, root='../data'): + """ + 加载 MNIST 数据集。 + + Parameters: + - data_size (int): 加载的数据条目数。 + - batch_size (int): 每批次处理的数据数量。 + - seed (int): 随机种子,用于保证数据切割的可重复性。 + - root (str): 数据集保存的路径。 + + Returns: + - train_loader (DataLoader): 训练数据集加载器。 + """ + transform = get_transform() + + 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 + ) + + return train_loader + + +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): 数据集保存的路径。 + """ + train_loader = load_mnist(data_size=data_size, batch_size=batch_size, seed=seed, root=root) + + 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) + mean = [0.5, 0.5, 0.5] + std = [0.5, 0.5, 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() + + # Or use custom parameters: + # load_and_preview_mnist(data_size=200, batch_size=8, seed=123, root='../data')