如何调整EMNIST数据集大小同时保持各类别样本占比平衡
实现类别平衡的PyTorch数据集调整方案
核心逻辑是先统计每个类别的样本索引,按目标数据集大小等比例分配每个类别的采样数量,均匀采样后合并索引封装为新数据集,天然保证类别占比和原数据集一致。
前置代码修正
你原代码存在splits变量定义顺序错误、单通道数据归一化参数不匹配的问题,先修正基础加载代码:
import torch import numpy as np from torch.utils.data import Subset from torchvision.datasets import EMNIST from torchvision import transforms from collections import Counter # 先定义split列表再调用 splits = ('byclass', 'bymerge', 'balanced', 'letters', 'digits', 'mnist') affine = transforms.RandomAffine([-15, 15], scale=(0.8, 1.2)) # 旋转与缩放 normalize = transforms.Normalize((0.0,), (1.0,)) # EMNIST为单通道灰度图,修正归一化参数 to_tensor = transforms.ToTensor() transform_train = transforms.Compose([to_tensor, affine]) # 加载原始数据集 EMNIST_train = EMNIST(root='./EMNIST_1st', split=splits[-2], train=True, download=True, transform=transform_train) EMNIST_test = EMNIST(root='./EMNIST_1st', split=splits[-2], train=False, download=True, transform=to_tensor)
平衡采样核心实现
def get_balanced_subset(dataset, target_size): # 获取全量标签 labels = np.array(dataset.targets) # 按类别分组存储样本索引 class_indices = {} for idx, label in enumerate(labels): if label not in class_indices: class_indices[label] = [] class_indices[label].append(idx) num_classes = len(class_indices) # 计算每类基础采样数 per_class_num = target_size // num_classes # 处理除不尽的余数,随机分配到不同类别避免总大小偏差 remainder = target_size % num_classes selected_indices = [] for cls, indices in class_indices.items(): # 打乱当前类索引避免采样偏差 np.random.shuffle(indices) take_num = per_class_num + 1 if remainder > 0 else per_class_num remainder -= 1 # 如果需要扩充数据集(目标大小大于原数据集),将下方切片改为 np.random.choice(indices, size=take_num, replace=True) 开启重复采样 selected_indices.extend(indices[:take_num]) # 打乱最终索引,避免同类别样本连续出现影响训练 np.random.shuffle(selected_indices) return Subset(dataset, selected_indices)
调用与验证示例
# 示例:训练集从240000缩到24000,测试集从40000缩到4000,保持类别平衡 balanced_train = get_balanced_subset(EMNIST_train, target_size=24000) balanced_test = get_balanced_subset(EMNIST_test, target_size=4000) # 验证数据集大小 print(f"调整后训练集大小:{len(balanced_train)},测试集大小:{len(balanced_test)}") # 验证类别平衡性 train_label_count = Counter([EMNIST_train.targets[i] for i in balanced_train.indices]) print(f"训练集每类样本数:{train_label_count}")
补充说明
- 若需要每类样本数量完全相等,可直接固定
per_class_num为指定值,不需要按总目标大小计算 - 可在采样前添加
np.random.seed(42)固定随机种子,保证采样结果可复现 - 若要扩大数据集,开启采样时的
replace=True参数即可实现类别平衡的过采样
内容的提问来源于stack exchange,提问作者unser
相关产品推荐
相关产品推荐

