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

如何平衡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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 17:20:24