如何为PyTorch Dataset/Dataloader使用平衡采样器实现指定批次比例
PyTorch固定比例正负样本的DataLoader平衡采样实现
需求回顾
需要实现每个训练批次由10个正样本 + 90个随机负样本组成,正样本数量不足时允许重复采样,不使用数据增强扩充样本,贴合PyTorch原生框架风格。
最优方案:自定义BatchBalancedSampler
PyTorch的Sampler抽象类是控制DataLoader采样逻辑的原生方式,相比WeightedRandomSampler,自定义Sampler能精准控制每个批次的正负样本比例,完全满足需求。
实现代码
import torch from torch.utils.data import Sampler import random from typing import Iterator class BatchBalancedSampler(Sampler[int]): def __init__(self, positive_idx: list, negative_idx: list, batch_size: int = 100, pos_per_batch: int = 10): self.positive_idx = positive_idx self.negative_idx = negative_idx self.batch_size = batch_size self.pos_per_batch = pos_per_batch self.neg_per_batch = batch_size - pos_per_batch # 对齐Dataset的总样本数 self.total_samples = 10000 self.total_batches = self.total_samples // self.batch_size def __iter__(self) -> Iterator[int]: sampler_indices = [] for _ in range(self.total_batches): # 正样本有放回采样:数量不足时自动重复选取 pos_samples = random.choices(self.positive_idx, k=self.pos_per_batch) # 负样本无放回采样:基于正负样本比例,负样本数量足够支撑无放回选取 neg_samples = random.sample(self.negative_idx, k=self.neg_per_batch) # 打乱批次内样本顺序,避免固定位置规律 batch = pos_samples + neg_samples random.shuffle(batch) sampler_indices.extend(batch) return iter(sampler_indices) def __len__(self) -> int: return self.total_samples
用法示例
ds = MyDataset() # 初始化自定义采样器 sampler = BatchBalancedSampler(ds.positive_idx, ds.negative_idx) # 配置DataLoader:关闭shuffle,使用自定义采样器 dl = torch.utils.data.DataLoader( ds, batch_size=100, shuffle=False, sampler=sampler )
关键逻辑说明
- 正样本重复采样:使用
random.choices实现有放回采样,当正样本数量小于10时,自动重复选取现有正样本,严格保证批次比例。 - 负样本采样:使用
random.sample实现无放回采样,结合用户给出的1:10000正负比例,负样本数量远大于90,完全满足无放回选取需求。 - 批次内打乱:每个批次的正负样本混合后打乱,避免模型学习到固定位置的样本规律。
- 原生框架兼容:继承PyTorch原生
Sampler类,无需修改原有Dataset逻辑,完全贴合PyTorch的设计风格。
极端情况处理
若负样本数量偶尔小于90(极端场景),可将random.sample替换为random.choices,改为负样本有放回采样,保证批次完整性。
内容的提问来源于stack exchange,提问作者Mateusz Konopelski
相关产品推荐
相关产品推荐

