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

如何为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
)

关键逻辑说明

  1. 正样本重复采样:使用random.choices实现有放回采样,当正样本数量小于10时,自动重复选取现有正样本,严格保证批次比例。
  2. 负样本采样:使用random.sample实现无放回采样,结合用户给出的1:10000正负比例,负样本数量远大于90,完全满足无放回选取需求。
  3. 批次内打乱:每个批次的正负样本混合后打乱,避免模型学习到固定位置的样本规律。
  4. 原生框架兼容:继承PyTorch原生Sampler类,无需修改原有Dataset逻辑,完全贴合PyTorch的设计风格。

极端情况处理

若负样本数量偶尔小于90(极端场景),可将random.sample替换为random.choices,改为负样本有放回采样,保证批次完整性。

内容的提问来源于stack exchange,提问作者Mateusz Konopelski

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 00:47:36