整体说明
- 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
28from 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 Sizeshuffle:训练时建议设为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)
- 迭代返回的批次默认是
补充理解:sampler 与 shuffle、dataset 的关系
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
)
sampler 和 dataset 的协作机制(呼应第一问)
即使通过
sampler=RandomSampler(self.train_dataset)传入了采样器(注:采样器会持有dataset的引用),**DataLoader依然必须同时持有dataset** ,因为内部流程固定为:1
2
3for 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_size、shuffle、drop_last和sampler都将失效(互斥),因为批次逻辑已完全交由batch_sampler自定义
StatefulDataLoader 使用说明
torchdata.stateful_dataloader.StatefulDataLoader是torchdata库中提供的一个数据加载器,可用于替代DataLoader- TorchData 项目聚焦于对
torch.utils.data.DataLoader的迭代增强 ,StatefulDataLoader是其首个重要成果
- TorchData 项目聚焦于对
StatefulDataLoader的核心目标是解决标准DataLoader无法在训练中途 保存和恢复迭代状态的问题
StatefulDataLoader 与 DataLoader 的核心区别
最根本的区别:有无状态管理
- 标准
DataLoader是无状态的- 每次重新创建迭代器都会从头开始遍历数据
StatefulDataLoader通过新增的state_dict()和load_state_dict()方法,让数据加载器具备了可保存、可恢复的能力特性 DataLoaderStatefulDataLoaderstate_dict()方法不支持 支持 load_state_dict()方法不支持 支持 训练中途断点续训 不支持 支持 保存当前迭代进度 无法保存 可精确保存已产出的批次数
构造参数:仅多一个参数
StatefulDataLoader的所有参数都与torch.utils.data.DataLoader完全相同 ,仅新增了一个关键字参数:snapshot_every_n_steps:Optional[int],默认值为1- 定义每隔多少步将状态从工作进程同步到主 DataLoader
StatefulDataLoader 的状态保存与恢复机制
StatefulDataLoader 默认行为(开箱即用)
- 默认情况下,
StatefulDataLoader的状态 包含已产生的批次数(number of batches yielded) - 恢复时,
StatefulDataLoader会利用这个数量 “快进”(fast-forward) 采样器(map-style 数据集)或数据集(iterable-style 数据集),精确定位到中断时的位置
StatefulDataLoader自定义状态(高级用法)
- 如果自定义
Sampler或Dataset也实现了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中的默认RandomSampler和BatchSampler,当导入torchdata.stateful_dataloader时,它们会被自动打补丁,因此无需自定义采样器 即可获得状态管理能力
StatefulDataLoader 使用示例
StatefulDataLoader 初始化(与 DataLoader 几乎完全一致)
1
2
3
4
5
6
7
8
9
10from 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必须一致 :- 恢复时要求
StatefulDataLoader的num_workers与保存state_dict时相同
- 恢复时要求
- 多进程安全 :如果使用多个 worker 线程,
dataset必须是 线程安全 的 - 状态同步仅支持单机多进程 :
- 状态同步不能跨越分布式训练中的不同 rank
- 版本兼容性 :
torchdata与torch有版本对应关系,安装时需要根据 PyTorch 版本选择合适的torchdata版本
StatefulDataLoader对标准 PyTorch 组件做了自动补丁(patched) :- 如果使用的是
torch.utils.data中的默认RandomSampler和BatchSampler,当导入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 数量
- 整体可以表述为:total_global_step 等于 数据集中包含的 Micro Batch 数量 除以 一个 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,这很反直觉切不容易排查
- 当样本数量为