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')