本文共 1043 字,大约阅读时间需要 3 分钟。
dataset参数是PyTorch中用于读取数据的接口,常见的有torchvision.datasets.ImageFolder或自定义数据接口的输出。输出应为torch.utils.data.Dataset类的对象,或继承自该类的自定义类对象。
根据具体需求设置合适的批量大小。通常用于控制每次数据加载的数量,建议根据GPU内存和训练任务规模进行调整。
在训练过程中,常常会启用shuffle功能,以确保数据的随机性,避免数据集的训练集和验证集划分不一致的情况。
用于对不同数据样本进行封装处理,通常采用默认设置,即使用DataLoader内置的collate函数,除非有特殊的数据读取需求。
与batch_size、shuffle等参数互斥,建议根据需求选择使用,默认情况下通常不需要自定义。
与shuffle互斥,通常默认使用,用于指定数据的抽取方式,例如随机抽取或循环抽取。
指定导入数据的工作数,0表示数据导入在主进程执行,>0时可以通过多线程加速数据导入速度。
布尔值,若为True,则将数据加载到GPU的 pinned memory 中,属于数据预加载的优化设置。
设置数据读取的超时时间,超过该时间未读取到数据时会抛出异常。
from torch.utils.data import DataLoaderfrom torchvision.datasets import ImageFolder# 示例数据集加载dataset = ImageFolder('data/your_data_path')# 创建DataLoader实例dataloader = DataLoader( dataset, batch_size=32, shuffle=True, num_workers=4, pin_memory=True)# 获取数据批次for batch in dataloader: inputs, labels = batch # 数据加载和训练逻辑