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

求基于PyTorch的CIFAR10可变批量大小DataLoader实现方案

适用于CIFAR10的可变批量大小数据加载器实现(批量大小服从随机几何分布)

实现思路

通过自定义BatchSampler动态生成服从随机几何分布的批量大小,替代PyTorch默认的固定批量采样逻辑,将该采样器传入DataLoader即可实现每次迭代使用不同批量大小的需求。

完整代码实现

import torch
import torchvision
import torchvision.transforms as transforms
from torch.utils.data import BatchSampler, DataLoader
from torch.distributions import Geometric

class GeometricBatchSampler(BatchSampler):
    def __init__(self, sampler, p, min_batch_size=1, drop_last=False):
        """
        服从随机几何分布的批量采样器
        Args:
            sampler: 基础采样器(如RandomSampler或SequentialSampler)
            p: 几何分布的成功概率参数(0 < p ≤ 1)
            min_batch_size: 最小批量大小,避免生成过小的批量
            drop_last: 是否丢弃最后一批不足随机生成大小的数据
        """
        super().__init__(sampler, batch_size=1, drop_last=drop_last)
        self.p = p
        self.min_batch_size = min_batch_size
        self.geometric_dist = Geometric(torch.tensor([p], dtype=torch.float32))
        
    def __iter__(self):
        batch = []
        for idx in self.sampler:
            batch.append(idx)
            # 生成随机批量大小,确保不小于最小批量
            target_batch_size = self.geometric_dist.sample().item() + self.min_batch_size
            if len(batch) >= target_batch_size:
                yield batch[:target_batch_size]
                batch = batch[target_batch_size:]
        # 处理剩余数据
        if batch and not self.drop_last:
            yield batch
            
    def __len__(self):
        # 近似长度,因批量大小随机,仅作参考
        total_samples = len(self.sampler)
        expected_batch_size = (1 / self.p) + self.min_batch_size - 1
        return int(total_samples / expected_batch_size) + (0 if self.drop_last else 1)

# 数据预处理
transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])

# 加载CIFAR10训练集
trainset = torchvision.datasets.CIFAR10(
    root='./data', train=True, download=True, transform=transform
)

# 使用随机采样器+自定义几何批量采样器
random_sampler = torch.utils.data.RandomSampler(trainset)
geometric_batch_sampler = GeometricBatchSampler(
    sampler=random_sampler, p=0.2, min_batch_size=8, drop_last=False
)

# 构建数据加载器
train_loader = DataLoader(
    trainset, batch_sampler=geometric_batch_sampler, num_workers=2
)

# 测试数据加载器
print("测试可变批量大小的数据加载器:")
for i, (images, labels) in enumerate(train_loader):
    print(f"第{i+1}批,批量大小:{len(images)},数据形状:{images.shape}")
    if i == 5:  # 仅打印前6批
        break

代码说明

  1. GeometricBatchSampler类:

    • 继承自PyTorch的BatchSampler,核心逻辑在__iter__方法中
    • 用torch.distributions.Geometric生成随机批量大小,通过min_batch_size避免极端小批量
    • 遍历基础采样器索引,累积到目标批量大小后返回批次,最后处理剩余数据
  2. 数据加载器构建:

    • 使用RandomSampler实现随机采样(需顺序采样可替换为SequentialSampler)
    • 将自定义采样器传入DataLoader的batch_sampler参数,此时无需指定batch_size
  3. 参数调整:

    • p:几何分布成功概率,值越小,期望批量越大(期望大小为(1/p) + min_batch_size - 1)
    • min_batch_size:设置最小批量,防止出现过小批次
    • drop_last:控制是否丢弃最后一批不足目标大小的数据

内容的提问来源于stack exchange,提问作者Naseem Yehya

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 21:10:18