diff --git a/Lab-2/src/__init__.py b/Lab-2/src/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/Lab-2/src/data_load.py b/Lab-2/src/data_load.py new file mode 100644 index 0000000..c91fab7 --- /dev/null +++ b/Lab-2/src/data_load.py @@ -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') \ No newline at end of file