Python多线程与多进程的结果复现:如何修复随机种子?
多进程多线程环境下的结果复现问题
我的代码执行流程如下:
- 启动数据采集进程
- 启动模型测试进程
- 一个线程负责训练(从采集进程读取数据)
- 一个线程负责测试(从测试进程读取数据)
- 训练线程每执行一步,需等待测试线程完成一步
- 测试线程执行前,需等待训练步骤完成
我需要实现结果复现,但进程和线程中均存在随机性。已在每个进程和线程中设置随机种子,但每次运行结果仍不一致。我使用两个线程池,每个仅启动1个线程,了解线程的非确定性,但不确定是否能实现完全复现。
目前的核心问题:通过池的initializer参数已让线程和进程内部行为具有确定性,但进程写入数据的顺序因多进程调度的随机性而无法固定——有时是某个进程先读取队列并写入,有时是另一个进程,导致最终结果不一致。
下面是简化的可复现代码(MWE),需要程序每次运行输出完全一致:
import logging import traceback import torch from concurrent.futures import ThreadPoolExecutor from concurrent.futures import ProcessPoolExecutor from torch import multiprocessing as mp shandle = logging.StreamHandler() log = logging.getLogger('rl') log.propagate = False log.addHandler(shandle) log.setLevel(logging.INFO) def fix_seed(seed): torch.manual_seed(seed) torch.cuda.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.manual_seed(seed) torch.backends.cudnn.benchmark = False torch.backends.cudnn.deterministic = True def collect(id, queue, data): #log.info('Collect %i started ...', id) while True: idx = queue.get() if idx is None: break data[idx] = torch.rand(1) log.info(f'Collector {id} got idx {idx} and sampled {data[idx]}') queue.task_done() #log.info('Collect %i completed', id) def test(id, queue, data): #log.info('Test %i started ...', id) while True: idx = queue.get() if idx is None: break data[idx] = torch.rand(1) log.info(f'Tester {id} got idx {idx} and sampled {data[idx]}') queue.task_done() #log.info('Test %i completed', id) def run(): steps = 0 num_collect_procs = 3 num_test_procs = 2 max_steps = 10 data_collect = torch.zeros(num_collect_procs).share_memory_() data_test = torch.zeros(num_test_procs).share_memory_() ctx = mp.get_context('spawn') manager = mp.Manager() collect_queue = manager.JoinableQueue() test_queue = manager.JoinableQueue() train_test_queue = manager.JoinableQueue() collect_pool = ProcessPoolExecutor( num_collect_procs, mp_context=ctx, initializer=fix_seed, initargs=(1,) ) test_pool = ProcessPoolExecutor( num_test_procs, mp_context=ctx, initializer=fix_seed, initargs=(1,) ) for i in range(num_collect_procs): future = collect_pool.submit(collect, i, collect_queue, data_collect) for i in range(num_test_procs): future = test_pool.submit(test, i, test_queue, data_test) def run_train(): nonlocal steps #log.info('Training thread started ...') while steps < max_steps: train_test_queue.put(True) train_test_queue.join() for idx in range(num_collect_procs): collect_queue.put(idx) log.info('Training, %i %f', steps, data_collect.sum() + torch.rand(1)) collect_queue.join() steps += 1 #log.info('Training ended') for i in range(num_collect_procs): collect_queue.put(None) train_test_queue.put(None) def run_test(): nonlocal steps #log.info('Testing thread started ...') while steps < max_steps: status = train_test_queue.get() if status is None: break for idx in range(num_test_procs): test_queue.put(idx) log.info('Testing, %i %f', steps, data_test.sum() + torch.rand(1)) test_queue.join() train_test_queue.task_done() #log.info('Testing ended') for i in range(num_test_procs): test_queue.put(None) training_thread = ThreadPoolExecutor(1, initializer=fix_seed, initargs=(1,)) testing_thread = ThreadPoolExecutor(1, initializer=fix_seed, initargs=(1,)) training_thread.submit(run_train) testing_thread.submit(run_test) if __name__ == '__main__': run()
解决方案
要实现完全的结果复现,核心是消除进程间的调度随机性,让每个进程的执行顺序和时机完全可控。具体可以从以下几个方面调整:
1. 替换争抢式队列,改为按进程ID定向同步
当前的JoinableQueue会让多个进程争抢任务,导致执行顺序随机。可以用Event同步机制替代队列,主进程按固定顺序触发每个进程的任务,等待该进程完成后再触发下一个,确保进程执行顺序完全固定。
2. 为每个进程分配独立固定的随机种子
所有进程共用同一个种子时,即使内部行为确定,执行顺序的变化仍会导致最终数据的写入顺序不同。给每个进程分配唯一的固定种子(如base_seed + 进程ID),确保每个进程生成的随机数完全独立且固定。
3. 线程内随机数生成单独固定种子
训练和测试线程内的torch.rand(1)也需要单独设置固定种子,避免线程调度的微小差异影响结果。
修改后的完整可复现代码
import logging import torch from concurrent.futures import ThreadPoolExecutor, ProcessPoolExecutor from torch import multiprocessing as mp shandle = logging.StreamHandler() log = logging.getLogger('rl') log.propagate = False log.addHandler(shandle) log.setLevel(logging.INFO) def fix_seed(seed): torch.manual_seed(seed) torch.cuda.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.backends.cudnn.benchmark = False torch.backends.cudnn.deterministic = True def collect(id, start_event, done_event, data, seed): fix_seed(seed) while True: start_event.wait() start_event.clear() # 终止信号:data设为None时退出 if data is None: break data[id] = torch.rand(1) log.info(f'Collector {id} got idx {id} and sampled {data[id]}') done_event.set() def test(id, start_event, done_event, data, seed): fix_seed(seed) while True: start_event.wait() start_event.clear() if data is None: break data[id] = torch.rand(1) log.info(f'Tester {id} got idx {id} and sampled {data[id]}') done_event.set() def run(): steps = 0 num_collect_procs = 3 num_test_procs = 2 max_steps = 10 base_seed = 1 data_collect = torch.zeros(num_collect_procs).share_memory_() data_test = torch.zeros(num_test_procs).share_memory_() ctx = mp.get_context('spawn') # 为每个进程创建同步事件 collect_start_events = [ctx.Event() for _ in range(num_collect_procs)] collect_done_events = [ctx.Event() for _ in range(num_collect_procs)] test_start_events = [ctx.Event() for _ in range(num_test_procs)] test_done_events = [ctx.Event() for _ in range(num_test_procs)] collect_pool = ProcessPoolExecutor(num_collect_procs, mp_context=ctx) test_pool = ProcessPoolExecutor(num_test_procs, mp_context=ctx) # 启动采集进程,每个进程用独立固定种子 for i in range(num_collect_procs): collect_pool.submit( collect, i, collect_start_events[i], collect_done_events[i], data_collect, base_seed + i ) # 启动测试进程,每个进程用独立固定种子 for i in range(num_test_procs): test_pool.submit( test, i, test_start_events[i], test_done_events[i], data_test, base_seed + num_collect_procs + i ) def run_train(): nonlocal steps fix_seed(base_seed) manager = mp.Manager() train_test_queue = manager.JoinableQueue() while steps < max_steps: # 等待测试线程完成当前步骤 train_test_queue.put(True) train_test_queue.join() # 按固定顺序触发采集进程任务 for idx in range(num_collect_procs): collect_start_events[idx].set() collect_done_events[idx].wait() collect_done_events[idx].clear() # 训练步骤的随机数生成 train_rand = torch.rand(1) log.info('Training, %i %f', steps, data_collect.sum() + train_rand) steps += 1 # 终止采集进程 for idx in range(num_collect_procs): global data_collect data_collect = None collect_start_events[idx].set() train_test_queue.put(None) def run_test(): nonlocal steps fix_seed(base_seed + num_collect_procs + num_test_procs) manager = mp.Manager() train_test_queue = manager.JoinableQueue() while steps < max_steps: status = train_test_queue.get() if status is None: break # 按固定顺序触发测试进程任务 for idx in range(num_test_procs): test_start_events[idx].set() test_done_events[idx].wait() test_done_events[idx].clear() # 测试步骤的随机数生成 test_rand = torch.rand(1) log.info('Testing, %i %f', steps, data_test.sum() + test_rand) train_test_queue.task_done() # 终止测试进程 for idx in range(num_test_procs): global data_test data_test = None test_start_events[idx].set() # 启动训练和测试线程 training_thread = ThreadPoolExecutor(1) testing_thread = ThreadPoolExecutor(1) training_thread.submit(run_train) testing_thread.submit(run_test) if __name__ == '__main__': run()
关键说明
- 用
Event同步替代队列,强制进程按固定顺序执行,彻底消除进程争抢的随机性。 - 每个进程使用独立固定种子,确保即使执行顺序(已固定)变化,生成的随机数仍完全确定。
- 训练、测试线程单独设置固定种子,避免线程调度的微小差异影响结果。
内容的提问来源于stack exchange,提问作者Simon
相关产品推荐
相关产品推荐

