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

如何在PyTorch中从Cifar-100按类别抽取10%样本构建自定义训练集

PyTorch下CIFAR-100每类按比例抽取训练样本的高效实现

你完全不需要手动为每个类别单独定义变量存储索引,以下两种方法都可以自动适配任意类别数的数据集,代码量不会随类别数量增长:

方法1:通用字典分组实现(兼容性最好)

该方法不需要额外依赖,兼容所有PyTorch内置数据集以及自定义数据集:

import random
from torch.utils.data import Subset
from torchvision import datasets

# 加载CIFAR-100训练集(你原有加载逻辑不变)
trainset = datasets.CIFAR100(root='./data', train=True, download=True, transform=None)

class_indices = {}
# 直接读取数据集内置的标签列表,无需遍历每个样本读取标签
for idx, label in enumerate(trainset.targets):
    if label not in class_indices:
        class_indices[label] = []
    class_indices[label].append(idx)

selected_indices = []
sample_ratio = 0.1
for indices in class_indices.values():
    # 计算单类采样数量,max避免样本数过少时采样数为0
    take_count = max(1, int(len(indices) * sample_ratio))
    # 若需要随机采样而非取前10%,打开下面这行打乱索引
    # random.shuffle(indices)
    selected_indices.extend(indices[:take_count])

# 生成最终的小子集
small_trainset = Subset(trainset, selected_indices)

方法2:Numpy向量化实现(性能更高)

如果数据集规模较大,用Numpy的向量化操作可以进一步降低遍历耗时:

import numpy as np
from torch.utils.data import Subset
from torchvision import datasets

trainset = datasets.CIFAR100(root='./data', train=True, download=True, transform=None)
targets = np.array(trainset.targets)
sample_ratio = 0.1

selected_indices = []
for label in np.unique(targets):
    # 批量获取该类别所有样本的索引
    label_indices = np.where(targets == label)[0]
    # 若需要随机采样,用下面这行代替切片逻辑
    # selected = np.random.choice(label_indices, size=max(1, int(len(label_indices)*sample_ratio)), replace=False)
    selected = label_indices[:max(1, int(len(label_indices)*sample_ratio))]
    selected_indices.extend(selected.tolist())

small_trainset = Subset(trainset, selected_indices)

注意事项

  • 两种方法都可以直接复用在CIFAR-10、ImageNet等其他PyTorch数据集上,不需要修改核心逻辑
  • 如果要保证采样无偏差,建议打开随机打乱索引的逻辑,避免取到的样本都是原始数据集同一分区的内容

内容的提问来源于stack exchange,提问作者Zixuan Liu

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 23:09:03