自定义PyTorch Dataset类调用len方法报method类型无len()错误如何解决
错误原因
processed_file_names是你定义的类实例方法,直接写self.processed_file_names拿到的是方法对象本身,不是方法执行后返回的文件名列表。方法对象不支持len()计算长度,因此触发类型错误。
修复方案
- 调用
processed_file_names时添加括号,获取它返回的列表后再计算长度 - 额外注意:你定义的
get方法中如果直接用传入的idx拼接文件名,会出现索引不匹配问题:PyTorch Dataset的索引默认从0开始,你现有文件名是data_1.pt到data_3.pt,调用get(0)时会尝试读取不存在的data_0.pt,需要同步修正索引逻辑。
修正后的完整类代码如下:
class Potsdam_Microgrid_DS1(Dataset): def __init__(self, root, transform=None, pre_transform=None): super().__init__(root, transform, pre_transform) #List all processed data def processed_file_names(self): return ['data_0.pt', 'data_1.pt', 'data_2.pt'] # 改为0开头匹配索引 def len(self): return len(self.processed_file_names()) # 添加括号调用方法 def get(self, idx): data = torch.load(osp.join(self.processed_dir, 'data_{}.pt'.format(idx))) return data
如果你不想修改已有的文件名,也可以只调整get方法的拼接逻辑:
def get(self, idx): data = torch.load(osp.join(self.processed_dir, 'data_{}.pt'.format(idx+1))) # idx加1匹配1开头的文件名 return data
内容的提问来源于stack exchange,提问作者BangNguyen
相关产品推荐
相关产品推荐

