PyTorch——DataLoader的使用


整体说明

  • PyTorch 中,Global Step 的计算和实现可能会因为不同版本而发生有趣的现象
  • 比如:在一些场景会遇到一些奇怪的现象,相差一个 样本,且不在 Global Batch 的整数倍边界,但是 Global Step 增加了 1

DataLoader 核心使用说明

  • torch.utils.data.DataLoader 用于加载数据集,实现批量读取、多线程加载、数据 Shuffle等功能
  • 用于训练和评估

DataLoader 基本用法

  • 用法简单示例:

    1
    2
    3
    4
    5
    6
    7
    8
    9
    10
    11
    12
    13
    14
    15
    16
    17
    18
    19
    20
    21
    22
    23
    24
    25
    26
    27
    28
    from torch.utils.data import DataLoader, Dataset

    # 定义数据集(需实现 __len__ 和 __getitem__)
    class MyDataset(Dataset):
    def __init__(self, data):
    self.data = data
    def __len__(self): # 返回总数据量
    return len(self.data)
    def __getitem__(self, idx): # 返回单条数据(idx为索引)
    return self.data[idx]

    # 初始化 DataLoader
    dataset = MyDataset(data=[1,2,3,4,5,6,7])
    dataloader = DataLoader(
    dataset=dataset, # 传入数据集
    batch_size=2, # 每个批次的样本数(默认1)
    shuffle=True, # 每个epoch是否打乱数据(默认False),与 sampler 互斥
    sampler=None, # 自定义索引采样器(如 RandomSampler / DistributedSampler)
    batch_sampler=None, # 自定义批次采样器(与 batch_size/shuffle/drop_last/sampler 互斥)
    drop_last=False, # 是否丢弃最后不完整批次(默认False)
    num_workers=0, # 加载数据的线程数(默认0,主线程加载)
    pin_memory=False, # 是否将数据存入固定内存(加速 GPU 读取,默认False)
    collate_fn=None # 自定义批处理拼接函数,用于将多行数据合并为一个 Batch
    )

    # 迭代使用(返回tensor批次)
    for batch in dataloader:
    print(batch)
  • 关键参数说明:

    • dataset:必须传入的数据集对象,这个类需实现 __len__/__getitem__
    • batch_size:批次大小,控制每次返回的样本数,当存在 Global Batch Size 和 Micro Batch Size,一般是 Micro Batch Size
    • shuffle:训练时建议设为 True,验证/测试时设为 False, 默认为 None
      • 注:设为 True 时,每个 epoch 开始时会重新打乱数据
    • drop_last:数据量无法被batch_size整除时,是否丢弃最后不足一个批次的样本
    • num_workers:多线程加载数据,可根据 CPU 核心数调整
    • pin_memory:若使用 GPU 训练,设为 True 可减少数据拷贝耗时,默认为 False
  • 注意事项

    • 迭代返回的批次默认是 tensor 类型(若数据集返回 numpy 或 列表 类型的数据,会自动转换);
    • shuffle=True 会在每个 epoch 开始时会重新打乱数据,保证训练随机性;
    • len(dataloader) 表示单个 epoch 的批次数量(计算规则:ceil(总数据量/batch_size)floor,取决于 drop_last

补充理解:samplershuffledataset 的关系

shuffle 本质是 sampler 的语法糖
  • shuffle=True 等价于手动指定 sampler=torch.utils.data.RandomSampler(dataset)
  • **shuffle=False**(默认)等价于手动指定 sampler=torch.utils.data.SequentialSampler(dataset)
  • 两者互斥 :如果显式传入了 sampler 参数,就**不能再设置 shuffle**(反之亦然),否则 PyTorch 会抛出 ValueError
  • 常见的正确的参数组合逻辑如下:
    1
    2
    3
    4
    5
    6
    7
    8
    9
    10
    11
    12
    13
    14
    15
    16
    17
    18
    # 方式一:使用 shuffle(隐式 SequentialSampler/RandomSampler)
    train_dataloader = DataLoader(
    dataset=train_dataset,
    batch_size=32,
    shuffle=True, # 等价于 sampler=RandomSampler(train_dataset)
    drop_last=True,
    num_workers=4
    )

    # 方式二:显式传入 sampler(此时必须去掉 shuffle)
    train_sampler = RandomSampler(train_dataset) # 或者 DistributedSampler(train_dataset)
    train_dataloader = DataLoader(
    dataset=train_dataset, # 依然必须显式传入!
    batch_size=32,
    sampler=train_sampler, # 显式指定,shuffle 必须为 False(默认)
    drop_last=True,
    num_workers=4
    )
samplerdataset 的协作机制(呼应第一问)
  • 即使通过 sampler=RandomSampler(self.train_dataset) 传入了采样器(注:采样器会持有 dataset 的引用),**DataLoader 依然必须同时持有 dataset** ,因为内部流程固定为:

    1
    2
    3
    for index in sampler:          # sampler 只负责生成索引(如 [3, 0, 7, 2])
    data = dataset[index] # dataset 负责根据索引取出真实样本(图像、标签)
    batch_data.append(data)
  • 理解:

    • sampler 只输出整数索引 ,不触碰数据
    • dataset 只负责索引取值 ,不决定顺序
  • 传入的 train_sampler 即使内部引用了 train_dataset(比如 DistributedSampler 需要知道数据集长度),DataLoader 也不会“聪明地”自动提取它作为数据源,必须显式传入 dataset 才能完成取值操作

高级参数:batch_sampler 的说明
  • batch_sampler 更高级,它直接生成批次的索引列表(如 [[3,0], [7,2], ...]
  • 一旦使用了 batch_sampler,则 batch_sizeshuffledrop_lastsampler 都将失效(互斥),因为批次逻辑已完全交由 batch_sampler 自定义

StatefulDataLoader 使用说明

  • torchdata.stateful_dataloader.StatefulDataLoadertorchdata 库中提供的一个数据加载器,可用于替代 DataLoader
    • TorchData 项目聚焦于对 torch.utils.data.DataLoader迭代增强StatefulDataLoader 是其首个重要成果
  • StatefulDataLoader 的核心目标是解决标准 DataLoader 无法在训练中途 保存和恢复迭代状态的问题

StatefulDataLoader 与 DataLoader 的核心区别

最根本的区别:有无状态管理
  • 标准 DataLoader 是无状态的
    • 每次重新创建迭代器都会从头开始遍历数据
  • StatefulDataLoader 通过新增的 state_dict()load_state_dict() 方法,让数据加载器具备了可保存、可恢复的能力
    特性 DataLoader StatefulDataLoader
    state_dict() 方法 不支持 支持
    load_state_dict() 方法 不支持 支持
    训练中途断点续训 不支持 支持
    保存当前迭代进度 无法保存 可精确保存已产出的批次数
构造参数:仅多一个参数
  • StatefulDataLoader所有参数都与 torch.utils.data.DataLoader 完全相同 ,仅新增了一个关键字参数:
    • snapshot_every_n_stepsOptional[int],默认值为 1
      • 定义每隔多少步将状态从工作进程同步到主 DataLoader

StatefulDataLoader 的状态保存与恢复机制

StatefulDataLoader 默认行为(开箱即用)
  • 默认情况下,StatefulDataLoader 的状态 包含已产生的批次数(number of batches yielded)
  • 恢复时,StatefulDataLoader 会利用这个数量 “快进”(fast-forward) 采样器(map-style 数据集)或数据集(iterable-style 数据集),精确定位到中断时的位置
StatefulDataLoader自定义状态(高级用法)
  • 如果自定义 SamplerDataset 也实现了 state_dict() / load_state_dict() 方法,StatefulDataLoader 会在自己的状态保存/恢复过程中自动调用它们 ,实现更深层次的状态管理
  • 例如:
    • Sampler 的自定义状态 :如随机数生成器(RNG)的种子和当前状态
    • Dataset 的自定义状态 :如 worker 特有的变换状态(RNG 变换状态)

StatefulDataLoader 的多进程(num_workers > 0)状态同步

  • 底层自动处理StatefulDataLoader 在底层处理多进程 workers 之间的状态聚合和分发
  • 跨 rank 不支持 :状态同步不跨越分布式训练中的不同 rank(即不处理跨 GPU/节点的同步)
  • 恢复时的约束 :调用 load_state_dict() 时,要求恢复后的 num_workers 与保存时保持一致

StatefulDataLoader 与自定义 Sampler/Dataset 的协作

  • StatefulDataLoader 对标准 PyTorch 组件做了自动补丁(patched)
  • 如果使用的是 torch.utils.data 中的默认 RandomSamplerBatchSampler,当导入 torchdata.stateful_dataloader 时,它们会被自动打补丁,因此无需自定义采样器 即可获得状态管理能力

StatefulDataLoader 使用示例

  • StatefulDataLoader 初始化(与 DataLoader 几乎完全一致)

    1
    2
    3
    4
    5
    6
    7
    8
    9
    10
    from torchdata.stateful_dataloader import StatefulDataLoader

    # 创建方式与 DataLoader 完全相同
    train_loader = StatefulDataLoader(
    dataset=train_dataset,
    batch_size=32,
    shuffle=True,
    num_workers=4,
    snapshot_every_n_steps=10 # 每10步同步一次worker状态
    )
  • StatefulDataLoader 保存和恢复状态

    1
    2
    3
    4
    5
    6
    7
    8
    9
    10
    11
    12
    13
    14
    15
    16
    17
    18
    19
    20
    21
    22
    23
    24
    25
    26
    27
    28
    29
    30
    31
    32
    33
    34
    35
    # ========== 训练过程中保存状态 ==========
    for i, batch in enumerate(train_loader):
    # ... 训练逻辑 ...

    if i == checkpoint_step:
    # 保存 DataLoader 状态
    dataloader_state = train_loader.state_dict()
    checkpoint = {
    'model': model.state_dict(),
    'optimizer': optimizer.state_dict(),
    'dataloader': dataloader_state,
    'epoch': epoch,
    'step': i
    }
    torch.save(checkpoint, 'checkpoint.pt')
    break

    # ========== 恢复训练时加载状态 ==========
    checkpoint = torch.load('checkpoint.pt')

    # 重新创建 DataLoader(必须与保存时 num_workers 一致)
    train_loader = StatefulDataLoader(
    dataset=train_dataset,
    batch_size=32,
    shuffle=True,
    num_workers=4 # 必须与保存时一致!
    )

    # 恢复 DataLoader 状态
    train_loader.load_state_dict(checkpoint['dataloader'])

    # 从中断处继续迭代
    for i, batch in enumerate(train_loader):
    # 从保存时的下一个 batch 开始继续训练
    # ...

补充:StatefulDataLoader 使用注意事项

  • num_workers 必须一致
    • 恢复时要求 StatefulDataLoadernum_workers 与保存 state_dict 时相同
  • 多进程安全 :如果使用多个 worker 线程,dataset 必须是 线程安全
  • 状态同步仅支持单机多进程
    • 状态同步不能跨越分布式训练中的不同 rank
  • 版本兼容性
    • torchdatatorch 有版本对应关系,安装时需要根据 PyTorch 版本选择合适的 torchdata 版本
  • StatefulDataLoader 对标准 PyTorch 组件做了自动补丁(patched)
    • 如果使用的是 torch.utils.data 中的默认 RandomSamplerBatchSampler,当导入 torchdata.stateful_dataloader 时,它们会被自动打补丁,因此无需自定义采样器 即可获得状态管理能力

附录:记一次有趣的 Bug 排查

  • 一般的 Global Batch Step 计算实现如下:

    1
    total_global_step = len(dataloader) * epoches // (global_batch_size / micro_batch_size)
    • 整体可以表述为:total_global_step 等于 数据集中包含的 Micro Batch 数量 除以 一个 Global Batch 中包含的 Micro Batch 数量
      • len(dataloader):数据集中包含的 Micro Batch 数量
      • (global_batch_size / micro_batch_size):一个 Global Batch 中包含的 Micro Batch 数量
  • 有趣的问题:假定 drop_last=False, micro_batch_size = 8, global_batch_size = 128, epoch=1

    • 总样本数量为 2040 时,total_global_step = 15
    • 总样本数量为 2041 时,total_global_step = 16
    • 可通过带入上面的代码验证结果
  • 问题出现的原因:

    • 当样本数量为 2040 时,len(dataloader) = 255
    • 当样本数量为 2041 时,len(dataloader) = 256
    • 进一步计算即可发现,虽然只是一个样本,且该数据量并不在 global_batch_size = 128 的整除边界,但是造成了 Global Step 多了 1,这很反直觉切不容易排查