PyTorch中DataLoader如何从Dataset获取并组装批次数据?
PyTorch DataLoader 对比学习批次组装原理解析
复现场景
使用PyTorch实现自监督对比学习时,自定义多视图增强变换代码如下:
class ContrastiveTransformations: def __init__(self, base_transforms, n_views=2): self.base_transforms = base_transforms self.n_views = n_views def __call__(self, x): return [self.base_transforms(x) for i in range(self.n_views)] contrast_transforms = transforms.Compose( [ transforms.RandomResizedCrop(size=96), transforms.ToTensor(), ] ) data_set = CIFAR10( root='/home1/data', download=True, transform=ContrastiveTransformations(contrast_transforms, n_views=2), )
单个数据集样本返回值为包含两个增强张量的列表[x_1, x_2],预期DataLoader返回的批次中data_batch是长度为batch_size的列表,每个元素对应单样本的[x_1, x_2],但实际输出格式为[[batch_x1, batch_x2], label_batch],即两个视图分别被组装为独立的批次张量。
查看DataLoader源码可知,map式数据集的取数逻辑为:
class _MapDatasetFetcher(_BaseDatasetFetcher): def __init__(self, dataset, auto_collation, collate_fn, drop_last): super(_MapDatasetFetcher, self).__init__(dataset, auto_collation, collate_fn, drop_last) def fetch(self, possibly_batched_index): if self.auto_collation: data = [self.dataset[idx] for idx in possibly_batched_index] else: data = self.dataset[possibly_batched_index] return self.collate_fn(data)
出现上述格式差异的核心在默认collate_fn的处理逻辑。
核心逻辑:default_collate的递归组装规则
PyTorch DataLoader默认使用torch.utils.data.default_collate做批次组装,它不会简单把取到的样本列表直接返回,而是递归遍历样本的嵌套结构,对同位置、类型/形状兼容的元素沿batch维度做拼接,处理流程如下:
- 取数阶段拿到的
data是长度为batch_size的列表,每个元素是CIFAR10返回的二元组:([x1_i, x2_i], y_i),其中i为批次内的样本序号。 - collate_fn首先识别到所有元素都是长度为2的元组,会拆分位置分别处理:
- 对元组第二位(标签位):收集所有样本的
y_i,这些值都是整数标量,直接拼接为形状为[batch_size]的标签张量label_batch。 - 对元组第一位(视图位):识别到每个样本的这个位置都是长度为2的列表,会继续拆分列表内的位置递归处理:
- 收集所有样本列表的第0个元素
x1_i,每个x1_i都是形状为[3, 96, 96]的图像张量,沿第0维拼接为形状[batch_size, 3, 96, 96]的batch_x1 - 收集所有样本列表的第1个元素
x2_i,同理拼接为同形状的batch_x2 - 两个拼接后的张量按照原样本的列表结构,组装为
[batch_x1, batch_x2]
- 收集所有样本列表的第0个元素
- 对元组第二位(标签位):收集所有样本的
- 两部分处理完成后,按照原样本的元组结构组装为最终返回值:
([batch_x1, batch_x2], label_batch),和实际运行观察到的格式完全一致。
这种递归拼接逻辑不需要额外手动处理视图的批次组装,天然适配对比学习的infoNCE损失计算输入要求。
自定义调整方法
如果需要其他批次格式,可以自定义collate_fn传入DataLoader构造函数,覆盖默认的递归拼接逻辑即可。
简单验证代码如下:
import torch from torch.utils.data import DataLoader from torchvision import transforms from torchvision.datasets import CIFAR10 # 运行后可直接观察到批次格式和描述一致 dataset = CIFAR10( root='./data', download=True, transform=ContrastiveTransformations(transforms.Compose([ transforms.RandomResizedCrop(96), transforms.ToTensor() ]), n_views=2) ) loader = DataLoader(dataset, batch_size=4) views, labels = next(iter(loader)) print(f"视图列表长度: {len(views)}") print(f"x1批次形状: {views[0].shape}") print(f"x2批次形状: {views[1].shape}") print(f"标签批次形状: {labels.shape}")
内容的提问来源于stack exchange,提问作者Peking Duck
相关产品推荐
相关产品推荐

