PyTorch 如何为自定义数据集手动应用归一化变换供DataLoader使用
自定义PyTorch数据集应用transform预处理方案
要实现和内置数据集完全一致的预处理逻辑,你只需要继承PyTorch的Dataset抽象类实现自定义数据集,将transform逻辑嵌入到样本读取流程中即可,具体实现如下:
首先导入依赖:
import torch import pandas as pd from torchvision import transforms from torch.utils.data import Dataset, DataLoader
然后实现自定义数据集类:
class SignMNISTDataset(Dataset): def __init__(self, csv_path, transform=None): self.data = pd.read_csv(csv_path) self.transform = transform def __len__(self): # 返回数据集总样本量 return len(self.data) def __getitem__(self, idx): # 拆分单条样本的标签和像素数据 label = self.data.iloc[idx, 0] # 784维像素值还原为28*28的单通道灰度图格式 img = self.data.iloc[idx, 1:].values.reshape(28, 28).astype('uint8') # 应用传入的预处理规则 if self.transform: img = self.transform(img) return img, label
之后就可以和教程中的内置数据集用法完全一致调用:
# 和教程保持完全相同的预处理规则 mean, std = (0.5,), (0.5,) transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean, std) ]) # 实例化自定义训练集 trainset = SignMNISTDataset( csv_path='../input/sign-language-mnist/sign_mnist_train.csv', transform=transform ) # 构造DataLoader trainloader = DataLoader(trainset, batch_size=64, shuffle=True)
以上实现的预处理逻辑和教程完全等价:ToTensor()会把0-255范围的uint8像素值转换为0-1范围的浮点张量,再经过Normalize处理后缩放至-1到1范围,适配后续神经网络输入要求。
内容的提问来源于stack exchange,提问作者kdeez
相关产品推荐
相关产品推荐

