求基于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
代码说明
GeometricBatchSampler类:
- 继承自PyTorch的
BatchSampler,核心逻辑在__iter__方法中 - 用
torch.distributions.Geometric生成随机批量大小,通过min_batch_size避免极端小批量 - 遍历基础采样器索引,累积到目标批量大小后返回批次,最后处理剩余数据
- 继承自PyTorch的
数据加载器构建:
- 使用
RandomSampler实现随机采样(需顺序采样可替换为SequentialSampler) - 将自定义采样器传入
DataLoader的batch_sampler参数,此时无需指定batch_size
- 使用
参数调整:
p:几何分布成功概率,值越小,期望批量越大(期望大小为(1/p) + min_batch_size - 1)min_batch_size:设置最小批量,防止出现过小批次drop_last:控制是否丢弃最后一批不足目标大小的数据
内容的提问来源于stack exchange,提问作者Naseem Yehya
相关产品推荐
相关产品推荐

