You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.04 10:45:28