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

如何查看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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 19:20:02