PyTorch中如何拼接ImageFolder数据集并保留classes属性?
解决ConcatDataset拼接后丢失classes属性的问题
这个问题很常见——ConcatDataset本身并没有继承单个数据集的classes属性,它的核心功能只是把多个数据集的样本序列拼接在一起,所以直接访问会抛出AttributeError。针对你的场景,有几种简便的解决方式:
最直接的方案(适配你的场景)
因为你的两个数据集train_set1_original和train_val_modified都是从同一个ImageFolder目录生成的,它们的classes属性完全一致,所以直接从任意一个原始ImageFolder实例中获取即可:
train_val_folder = root + "train_val_data/" train_set1_original = datasets.ImageFolder(train_val_folder, transform=train_val_transform_croponly) train_val_modified = datasets.ImageFolder(train_val_folder, transform=train_val_transform) train_val_dataset = torch.utils.data.ConcatDataset([train_set1_original, train_val_modified]) # 直接从原始ImageFolder数据集提取classes classes = train_set1_original.classes # 用train_val_modified.classes结果完全相同 train_val_length = len(train_val_dataset) print(train_val_length) print(classes) # 现在可以正常输出类别列表了
可选:给ConcatDataset实例手动绑定属性
如果你希望后续能直接通过train_val_dataset.classes访问类别信息,也可以手动给拼接后的数据集实例添加这个属性:
train_val_dataset = torch.utils.data.ConcatDataset([train_set1_original, train_val_modified]) # 手动添加classes和class_to_idx属性 train_val_dataset.classes = train_set1_original.classes train_val_dataset.class_to_idx = train_set1_original.class_to_idx # 之后就能直接用了 print(train_val_dataset.classes)
通用方案(适合多数据集拼接场景)
如果以后需要拼接不同来源但类别体系一致的数据集,或者想让拼接后的数据集自动具备classes属性,可以继承ConcatDataset写一个轻量级子类:
class ConcatDatasetWithClasses(torch.utils.data.ConcatDataset): def __init__(self, datasets): super().__init__(datasets) # 假设所有输入数据集的类别体系一致,取第一个数据集的属性即可 self.classes = datasets[0].classes self.class_to_idx = datasets[0].class_to_idx # 使用方式和原生ConcatDataset完全一致 train_val_dataset = ConcatDatasetWithClasses([train_set1_original, train_val_modified]) print(train_val_dataset.classes) # 正常访问类别
对你的场景来说,第一种方案是最省事高效的,毕竟两个数据集来自同一目录,类别信息完全同步,直接复用原始实例的属性就好。
内容的提问来源于stack exchange,提问作者roundboxes
相关产品推荐
相关产品推荐

