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

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]
  • 两部分处理完成后,按照原样本的元组结构组装为最终返回值:([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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 12:31:09