PyTorch初学者加载NYU数据集时触发pic应为Tensor/ndarray的NoneType报错
报错原因排查
- 核心原因是类方法缩进错误:你编写的
__getitem__、__len__两个方法没有缩进,不属于Mydataset类的成员方法。PyTorch的Dataset子类必须重载这两个方法才能正常读取数据,你没有正确重载的情况下,DataLoader调用数据读取接口时只会返回None,对应你收到的NoneType报错。 - 数据变换逻辑存在错误:你定义的transform第一个操作是
transforms.ToPILImage(),但通过Image.open()读取的输入本身就是PIL格式图像,ToPILImage()仅支持Tensor、ndarray类型的输入,传入PIL对象会触发类型不匹配问题。 - Resize传参不符合规范:
transforms.Resize()如果要指定输出的高宽,需要传入(height, width)格式的元组,你写的Resize(224,101)会把224识别为等比缩放的短边长度,101被错误识别为插值方法参数,无法得到你预期的输出尺寸。
解决方法
- 修正类方法缩进,将
__getitem__、__len__两个方法整体缩进,纳入Mydataset类的代码块 - 删除transform流水线中的
transforms.ToPILImage()操作 - 修正Resize的传参格式,改为传入元组指定输出尺寸
修正后的完整代码如下:
from __future__ import division,absolute_import,print_function from PIL import Image from torch.utils.data import DataLoader,Dataset from torchvision.transforms import transforms data_root='D:/AuxiliaryDocuments/NYU/' transforms=transforms.Compose([ transforms.Resize((224,101)), transforms.ToTensor()]) filename_txt={'image_train':'image_train.txt','image_test':'image_test.txt', 'depth_train':'depth_train.txt','depth_test':'depth_test.txt'} class Mydataset(Dataset): def __init__(self,data_root,transformation,data_type): self.transform=transformation self.image_path_txt=filename_txt[data_type] self.sample_list=list() f=open(data_root+'/'+data_type+'/'+self.image_path_txt) lines=f.readlines() for line in lines: line=line.strip() line=line.replace(';','') self.sample_list.append(line) f.close() def __getitem__(self, index): item=self.sample_list[index] img=Image.open(item) if self.transform is not None: img=self.transform(img) idx=index return idx,img def __len__(self): return len(self.sample_list)
内容的提问来源于stack exchange,提问作者Feona
相关产品推荐
相关产品推荐

