如何平衡PyTorch不平衡数据集?支持删除多数类样本
不平衡PyTorch数据集的下采样平衡方法
针对你需要通过删除多数类样本实现数据集平衡的需求,以下是几种更便捷的实现方式,涵盖sklearn和PyTorch原生方案:
一、使用sklearn的resample工具
利用sklearn.utils.resample可以快速对每个类别进行下采样,统一到最小类别(A/V类)的样本数量,步骤如下:
import torch from sklearn.utils import resample from torch.utils.data import Subset # 替换为你的自定义数据集实例 dataset = YourCustomDataset(...) # 1. 按类别分组样本索引 class_indices = {} for idx, (_, label) in enumerate(dataset): # 兼容张量和数值类型的标签 label = label.item() if isinstance(label, torch.Tensor) else label if label not in class_indices: class_indices[label] = [] class_indices[label].append(idx) # 2. 确定目标样本数(取样本最少的类别数量) target_sample_count = min(len(indices) for indices in class_indices.values()) # 3. 对每个类别随机下采样到目标数量 balanced_indices = [] for indices in class_indices.values(): # replace=False表示不重复采样(即删除多余样本) sampled_indices = resample(indices, n_samples=target_sample_count, random_state=42, replace=False) balanced_indices.extend(sampled_indices) # 4. 生成平衡后的数据集 balanced_dataset = Subset(dataset, balanced_indices)
二、PyTorch原生自定义实现(无需依赖sklearn)
如果不想引入额外依赖,手动实现随机下采样同样简单:
import torch import random from torch.utils.data import Subset dataset = YourCustomDataset(...) # 按类别分组索引 class_indices = {} for idx, (_, label) in enumerate(dataset): label = label.item() if isinstance(label, torch.Tensor) else label class_indices.setdefault(label, []).append(idx) target_sample_count = min(len(v) for v in class_indices.values()) balanced_indices = [] for indices in class_indices.values(): # 随机打乱后截取前target_sample_count个样本 random.shuffle(indices) balanced_indices.extend(indices[:target_sample_count]) balanced_dataset = Subset(dataset, balanced_indices)
三、动态采样(可选)
如果不想提前删除样本,而是在DataLoader层面动态平衡,可以使用WeightedRandomSampler,但这种方式是动态调整采样概率,不会实际删除样本:
import torch from torch.utils.data import WeightedRandomSampler dataset = YourCustomDataset(...) # 计算每个样本的权重(与类别样本数成反比) class_counts = {} for _, label in dataset: label = label.item() if isinstance(label, torch.Tensor) else label class_counts[label] = class_counts.get(label, 0) + 1 sample_weights = [] for _, label in dataset: label = label.item() if isinstance(label, torch.Tensor) else label sample_weights.append(1.0 / class_counts[label]) sampler = WeightedRandomSampler(sample_weights, num_samples=len(dataset), replacement=True) # 使用该sampler创建DataLoader balanced_loader = torch.utils.data.DataLoader(dataset, sampler=sampler, batch_size=32)
以上方法都比手动逐个过滤样本更高效、易维护,且支持随机种子控制保证实验可复现。
内容的提问来源于stack exchange,提问作者Yaroslav
相关产品推荐
相关产品推荐

