如何在PyTorch自定义Dataset上应用自定义Transform数据变换
PyTorch自定义数据集Transform不生效问题解决方案
核心原因
- 你打印的是数据集实例的
samples属性,该属性是你初始化数据集时传入的原始numpy数组,Transform的变换逻辑仅在__getitem__方法中触发,只有通过索引访问数据集元素(比如transformed[0])、或者使用DataLoader加载数据时才会执行变换,直接访问原始存储的samples属性不会触发任何变换操作。 - 你传入数据集的原始数据是numpy数组类型,但自定义的Transform都是基于PyTorch张量的操作,直接传入numpy数组会导致类型不匹配,需要先把numpy数组转成PyTorch张量再执行变换。
- 测试代码中你定义的数据集类名为
MyDataset,但实例化时用的是NumbersDataset,类名不一致会触发命名错误,属于笔误类问题。
修复后的可运行代码示例
import torch import numpy as np from torch.utils.data import Dataset from torchvision import transforms # 自定义高斯噪声变换 class AddGaussianNoise(object): def __init__(self, mean, std): self.std = std self.mean = mean def __call__(self, tensor): return tensor + torch.randn(tensor.size()) * self.std + self.mean def __repr__(self): return self.__class__.__name__ + f'(mean={self.mean}, std={self.std})' # 自定义归一化变换 class Normalize(object): def __init__(self, mean, std): self.std = std self.mean = mean def __call__(self, tensor): return (tensor.sub_(self.mean)).div(self.std) def __repr__(self): return self.__class__.__name__ + f'(mean={self.mean}, std={self.std})' # 自定义数据集 class MyDataset(Dataset): def __init__(self, data, transforms = None): self.samples = data self.transforms= transforms def __len__(self): return len(self.samples) def __getitem__(self, idx): sample = self.samples[idx] # 先把numpy数组转换为float类型的张量,适配后续变换操作 sample = torch.tensor(sample, dtype=torch.float32) if self.transforms is not None: sample = self.transforms(sample) return sample # 测试逻辑 if __name__ == "__main__": data = np.array([[-1,-1,1,-1],[-1,1,-1,-1],[1,-1,-1,-1],[-1,-1,-1,1]]) transformed_dataset = MyDataset(data, transforms.Compose([ AddGaussianNoise(0.5, 0.5), Normalize(0.5, 0.5), ])) # 通过索引访问元素触发变换,查看效果 for idx in range(len(transformed_dataset)): print(f"第{idx}个样本变换后结果:\n{transformed_dataset[idx]}\n")
补充说明
你单独调用Transform可以正常生效的原因是,你直接把完整数据传给了Transform实例主动执行了变换,没有经过数据集的懒加载逻辑,所以可以直接得到变换后的结果。如果你的需求是对整个二维数组做变换,而非按行拆分样本,修改__getitem__方法的索引逻辑即可。
内容的提问来源于stack exchange,提问作者samiogx
相关产品推荐
相关产品推荐

