使用PyTorch进行数据增强时变换未生效的问题咨询
问题分析与修复方案
1. 变换代码未执行的核心原因
你仅完成了数据集的实例化,但**__getitem__方法只有在通过索引访问数据集元素时才会触发**(比如执行dataset_pred_trans[0])。仅仅创建dataset_pred_trans对象不会运行任何数据加载和变换逻辑,自然不会执行if self.transform里的代码块。
2. 隐藏的类型不匹配问题
即使触发了__getitem__,你还会遇到报错:从h5py读取的image_amplitude和image_phase是numpy数组,但torchvision.transforms.Resize默认只支持PIL Image或torch.Tensor类型,直接传入numpy数组会抛出类型错误。需要补充类型转换步骤:
修改变换链:
transforms = torchvision.transforms.Compose([ torchvision.transforms.ToPILImage(), # 将numpy数组转为PIL Image torchvision.transforms.Resize((64,64)), torchvision.transforms.ToTensor() # 可选,转为Tensor适配PyTorch模型 ])
验证与调试步骤
- 实例化数据集后,手动访问元素触发逻辑:
sample = dataset_pred_trans[0] # 检查变换后的尺寸是否符合预期 print(sample['image_amplitude'].shape)
- 可在
__init__中添加打印,确认变换对象是否被正确传入:
def __init__(self, h5py_file, subset='train', transform=None): self.f = h5py.File(h5py_file,'r') self.transform = transform self.subset = subset print(f"传入的变换对象: {self.transform}") # 验证传参是否成功
内容的提问来源于stack exchange,提问作者Peter
相关产品推荐
相关产品推荐

