如何查看PyTorch Dataset子集的标签分布?
解决方法
当你用torch.utils.data.random_split()得到的是Subset对象,它本身没有targets属性——这个类只存储了原数据集和对应的索引列表,所以直接访问subset.targets会报错。你可以用下面两种方法统计标签分布:
方法1:遍历子集收集标签
直接遍历子集里的每个样本,提取标签后用Counter统计:
from collections import Counter # 假设subset是你拆分得到的子集 labels = [sample[1] for sample in subset] print(dict(Counter(labels)))
注:这里默认你的数据集每个样本是(数据, 标签)的元组结构,如果你的标签存储位置不同(比如样本是字典,标签存在sample['label']),对应调整提取方式即可。
方法2:通过原数据集的targets和子集索引获取
如果你的主数据集本身有targets属性(比如MNIST、CIFAR等Torch内置数据集),可以利用子集的indices属性获取对应索引的标签:
from collections import Counter import torch # original_dataset是你的主数据集,subset是拆分后的子集 if isinstance(original_dataset.targets, torch.Tensor): labels = original_dataset.targets[subset.indices].numpy() else: labels = [original_dataset.targets[i] for i in subset.indices] print(dict(Counter(labels)))
这里处理了targets是Tensor或列表的两种常见情况,确保Counter能正常统计。
内容的提问来源于stack exchange,提问作者Zowie Tay
相关产品推荐
相关产品推荐

