博客
关于我
PyTorch之DataLoader杂谈
阅读量:798 次
发布时间:2023-03-04

本文共 2859 字,大约阅读时间需要 9 分钟。

PyTorch 数据加载器(DataLoader)详解

在PyTorch中,数据加载器(DataLoader)是构建训练流程的核心模块之一。它负责将自定义数据集按照批量大小和其他配置参数封装成Tensor形式,从而为模型提供训练数据。作为数据进入模型的关键步骤,DataLoader的设计和实现至关重要。本文将深入剖析DataLoader的工作原理,包括其源码结构、参数配置以及与Sampler和BatchSampler的关系。

DataLoader 的核心作用

DataLoader的主要职责是根据数据集和批量配置(如batch_size、shuffle等),将数据以Tensor形式组织为小批量的数据样本。这些小批量数据被称为“Batch”,然后通过 DataLoader 加载到模型中进行训练。其工作流程一般是:

  • 创建一个 Dataset 对象
  • 创建一个 DataLoader 对象
  • 循环 DataLoader 对象,逐个加载 img 和 label 数据到模型中
  • DataLoader 的参数配置

    DataLoader 的初始化参数如下:

    • dataset:传入的数据集对象
    • batch_size(可选,默认为1):每个batch包含的样本数量
    • shuffle(可选,默认为False):每个epoch是否重新排序数据
    • sampler(可选):自定义数据采样策略,需确保 shuffle 为False
    • batch_sampler(可选):一次返回一个batch的索引数组
    • num_workers(可选,默认为0):数据加载时使用的工作进程数量
    • collate_fn(可选):将多个样本组成一个batch的函数
    • pin_memory(可选,默认为False):数据加载时是否使用GPU内存
    • drop_last(可选,默认为False):是否舍弃最后一个不完整的batch
    • timeout(可选,默认为0):等待数据加载的超时时间
    • worker_init_fn(可选):每个工作进程的初始化函数

    DataLoader 源码剖析

    DataLoader 类的实现主要集中在 inititer 方法中。以下是其核心代码:

    class DataLoader(object):
    __initialized = False
    def __init__(self, dataset, batch_size=1, shuffle=False, sampler=None,
    batch_sampler=None, num_workers=0, collate_fn=default_collate,
    pin_memory=False, drop_last=False, timeout=0,
    worker_init_fn=None):
    # 参数赋值和初始化检查
    # ...
    # 如果设置了 batch_sampler,需要确保与其他参数互斥
    # ...
    # 初始化 sampler 和 batch_sampler
    # ...
    self.__initialized = True
    def __setattr__(self, attr, val):
    # 禁止在初始化后设置某些属性
    # ...
    super(DataLoader, self).__setattr__(attr, val)
    def __iter__(self):
    return _DataLoaderIter(self)
    def __len__(self):
    return len(self.batch_sampler)

    DataLoader 的工作流程

  • Sampler 和 BatchSampler:DataLoader 的核心在于如何从数据集中采样。PyTorch 提供了三种主要的Sampler:

    • SequentialSampler:按顺序逐个采样
    • RandomSampler:随机采样,确保不重复
    • BatchSampler:基于Sampler获取批量索引

    BatchSampler 的作用是将Sampler的输出按批量形式组织。例如:

    # 示例:使用 SequentialSampler 和 BatchSampler
    sampler = SequentialSampler(range(10))
    batch_sampler = BatchSampler(sampler, batch_size=3, drop_last=False)
    dataloader = DataLoader(dataset=sampler, batch_sampler=batch_sampler)
  • DataLoaderIterator:DataLoader 的 __iter__ 方法返回一个 _DataLoaderIter 实例。该迭代器负责处理数据加载的逻辑,包括多进程和多线程的数据读取策略。

  • 多进程与多线程支持:DataLoader 可以根据 num_workers 参数配置多进程或多线程来加速数据加载。默认情况下,数据加载会在主进程中完成(num_workers=0)。

  • 数据读取与内存管理:如果 pin_memory=True,DataLoader 会在返回数据之前将其拷贝到GPU的固定内存中,提升数据加载速度。

  • DataLoader 的实际应用

    在实际训练中,DataLoader 的应用通常遵循以下流程:

    # 创建数据集
    dataset = MyDataset()
    # 创建 DataLoader
    dataloader = DataLoader(dataset, batch_size=32, shuffle=True, num_workers=4)
    # 训练过程
    num_epochs = 100
    for epoch in range(num_epochs):
    for img, label in dataloader:
    # img 和 label 是 tensors
    # 进行模型训练或推理
    pass

    总结

    DataLoader 是 PyTorch 中数据加载的核心工具,其功能涵盖了数据集的批量化处理、多进程加速、随机采样和内存管理等多个方面。通过合理配置 DataLoader 的参数,可以显著提升数据加载的效率和训练的稳定性。在实际应用中,理解和配置 DataLoader 的行为,是成为PyTorch高效训练模型的关键技能之一。

    转载自:Carson Zhu的博客

    你可能感兴趣的文章
    POJ 1088 滑雪
    查看>>
    POJ 1095 Trees Made to Order
    查看>>
    POJ 1113 Wall(计算几何--凸包的周长)
    查看>>
    poj 1125Stockbroker Grapevine(最短路)
    查看>>
    Qualitor processVariavel.php 未授权命令注入漏洞复现(CVE-2023-47253)
    查看>>
    poj 1151 (未完成) 扫描线 线段树 离散化
    查看>>
    POJ 1151 / HDU 1542 Atlantis 线段树求矩形面积并
    查看>>
    poj 1163 数塔
    查看>>
    POJ 1177 Picture(线段树:扫描线求轮廓周长)
    查看>>
    Qualitor checkAcesso.php 任意文件上传漏洞复现(CVE-2024-44849)
    查看>>
    POJ 1182 食物链(并查集拆点)
    查看>>
    POJ 1185 炮兵阵地 (状态压缩DP)
    查看>>
    POJ 1195 Mobile phones
    查看>>
    POJ 1228 Grandpa's Estate (稳定凸包)
    查看>>
    poj 1236(强连通分量分解模板题)
    查看>>
    poj 1258 Agri-Net
    查看>>
    quagga 和 zebos
    查看>>
    poj 1286 Necklace of Beads
    查看>>
    POJ 1321 棋盘问题
    查看>>
    poj 1321(回溯)
    查看>>