PyTorch DataLoader遍历pandas存路径的自定义数据集时出现KeyError怎么解决
报错根本原因
- pandas DataFrame 索引逻辑误用
你最初写的self.image_list[idx]对于pandas的DataFrame对象而言,[]运算符默认是按列名匹配取值,而非按行的位置取值。你的DataFrame仅存在一个名为path的列,没有任何整数类型的列名,因此传入整数类型的idx时,本质是在查询名为对应整数的列,自然触发KeyError。 - 报错key值随机的原因
你初始化DataLoader时设置了shuffle=True,DataLoader每次调用__getitem__方法时传入的idx是随机生成的采样序号,因此每次报错提示找不到的列名也就对应为随机的整数值。
修改后恢复正常的原因
你调用pd.read_csv时指定了index_col=False,pandas会自动为生成的DataFrame设置从0开始、连续递增的整数行索引,和DataLoader传入的、代表样本位置的idx刚好完全对齐。而loc运算符的作用是按行索引标签+列名取值,此时行标签与样本位置序号完全匹配,因此可以正确取到对应行的path列的值。
优化建议
- 如果后续你会对
image_list做过滤、去重、删除空值等修改操作,会导致行索引不再连续,此时用loc也可能触发KeyError,更稳妥的写法是用iloc按行的物理位置取值:self.image_list.iloc[idx]['path'],完全适配PyTorch Dataset的位置索引逻辑。 - 也可以在Dataset初始化阶段直接把路径转为普通列表存储,彻底规避pandas索引的坑:
self.image_list = pd.read_csv(csv_file, names=['path'], index_col=False)['path'].tolist(),后续__getitem__中直接用self.image_list[idx]即可像操作普通列表一样取值,逻辑更简单不易出错。
内容的提问来源于stack exchange,提问作者starc52
相关产品推荐
相关产品推荐

