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

Python多线程与多进程的结果复现:如何修复随机种子?

多进程多线程环境下的结果复现问题

我的代码执行流程如下:

  1. 启动数据采集进程
  2. 启动模型测试进程
  3. 一个线程负责训练(从采集进程读取数据)
  4. 一个线程负责测试(从测试进程读取数据)
  5. 训练线程每执行一步,需等待测试线程完成一步
  6. 测试线程执行前,需等待训练步骤完成

我需要实现结果复现,但进程和线程中均存在随机性。已在每个进程和线程中设置随机种子,但每次运行结果仍不一致。我使用两个线程池,每个仅启动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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 00:34:54