You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.19 09:34:54