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

如何调整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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 17:27:03