🔄Update: [Lab-2] 数据处理变换
This commit is contained in:
@@ -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')
|
||||||
Reference in New Issue
Block a user