PyTorch中如何保存恢复DataLoader的persistent_workers状态以续训
解决PyTorch persistent_workers=True时num_workers=0与>0训练结果不一致的问题
问题核心缺失
当前操作遗漏了两个关键环节:
- 未为worker进程设置独立且可复现的随机种子:你的
_init_fn是空实现,导致每个worker的随机生成器状态不可控,和主进程(num_workers=0时使用的)随机状态无法对齐。 - 未保存/恢复worker进程的随机生成器状态:当
persistent_workers=True时,worker进程会在epoch之间持续存在,它们的随机状态会随数据加载不断变化,但你只保存了主进程的随机状态,没有同步worker的状态。
具体实现方案
1. 修复worker_init_fn,为每个worker分配独立种子
每个worker需要有唯一的种子,确保不同num_workers配置下的随机性可复现。可以基于主种子+worker_id生成worker专属种子:
def _init_fn(worker_id): # 基于PyTorch自动分配的初始种子+worker_id生成唯一种子 seed = torch.initial_seed() % 2**32 seed += worker_id np.random.seed(seed) random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed(seed)
注:PyTorch的DataLoader会自动为每个worker生成初始种子(通过
torch.initial_seed()获取),在此基础上加上worker_id可确保每个worker的种子唯一且可复现。
2. 保存与恢复worker进程的随机状态
当persistent_workers=True时,需要在每个epoch结束时收集所有worker的随机状态并保存到检查点;恢复时将状态传递给对应的worker。
收集worker状态
使用多进程共享字典存储worker状态,配合自定义Dataset实现状态追踪:
import multiprocessing # 初始化共享字典,用于存储每个worker的随机状态 worker_states = multiprocessing.Manager().dict() def _init_fn(worker_id): # 初始化worker种子(同上) seed = torch.initial_seed() % 2**32 seed += worker_id np.random.seed(seed) random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed(seed) # 绑定worker_id与进程,方便后续收集状态 worker_states[worker_id] = {} # 自定义Dataset,在样本加载时更新worker状态 class StateTrackingDataset(Dataset): def __init__(self, original_dataset): self.original_dataset = original_dataset def __getitem__(self, idx): item = self.original_dataset[idx] worker_info = torch.utils.data.get_worker_info() if worker_info is not None: # 保存当前worker的随机状态 worker_states[worker_info.id] = { 'torch_rng_state': torch.get_rng_state(), 'numpy_rng_state': np.random.get_state(), 'py_rng_state': random.getstate() } return item def __len__(self): return len(self.original_dataset)
在每个epoch结束保存检查点时,将worker状态加入检查点:
checkpoint = { 'cpu_rng_state': torch.get_rng_state(), 'gpu_rng_state': torch.cuda.get_rng_state(), 'gpu_rng_state_all': torch.cuda.get_rng_state_all(), 'numpy_rng_state': np.random.get_state(), 'py_rng_state': random.getstate(), 'worker_rng_states': dict(worker_states) # 保存所有worker的状态 } torch.save(checkpoint, 'your_checkpoint.pth')
恢复worker状态
在恢复训练前,从检查点加载worker状态,并修改_init_fn实现状态恢复:
# 全局变量,存储从检查点加载的worker状态 restored_worker_states = {} def _init_fn(worker_id): if worker_id in restored_worker_states: # 恢复该worker的随机状态 state = restored_worker_states[worker_id] torch.set_rng_state(state['torch_rng_state']) np.random.set_state(state['numpy_rng_state']) random.setstate(state['py_rng_state']) else: # 首次初始化时设置种子 seed = torch.initial_seed() % 2**32 seed += worker_id np.random.seed(seed) random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed(seed) worker_states[worker_id] = {} # 恢复训练前加载检查点 checkpoint = torch.load('your_checkpoint.pth') restored_worker_states = checkpoint['worker_rng_states'] # 恢复主进程状态 torch.set_rng_state(checkpoint['cpu_rng_state']) torch.cuda.set_rng_state(checkpoint['gpu_rng_state']) torch.cuda.set_rng_state_all(checkpoint['gpu_rng_state_all']) np.random.set_state(checkpoint['numpy_rng_state']) random.setstate(checkpoint['py_rng_state'])
额外说明
- 当
persistent_workers=False时,worker会在每个epoch后重启,每次都会通过_init_fn重新初始化种子,因此不需要保存worker状态;但persistent_workers=True时worker持续存在,必须保存状态才能保证恢复后的一致性。 - 创建DataLoader时需指定
worker_init_fn=_init_fn,并将原始Dataset替换为StateTrackingDataset。
内容的提问来源于stack exchange,提问作者HATEM EL-AZAB
相关产品推荐
相关产品推荐

