如何取消numpy已播种的随机序列?多进程+@jitclass场景问题
我明白你现在的困境:想让多进程里第一部分随机数用统一固定种子(所有进程结果一致),第二部分要每个进程生成不一样的随机序列,但因为之前固定了种子,用np.random.randint生成新种子时,每个进程出的数都一样,导致后续序列还是重复,而且受限于@jitclass,不能用np.random.seed(None)或者时间戳种子来破局。
问题根源拆解
你看你的输出就知道问题在哪:第一组随机数完全一致(这是符合预期的),但第二组也一模一样——因为在统一种子的状态下,所有进程调用np.random.randint(100000000)时,生成的随机整数是完全相同的,相当于又给所有进程设了同一个新种子,自然后续序列也重复了。
靠谱的解决方案:用进程唯一标识做种子偏移
既然不能用时间或者None种子,我们可以用每个进程独有的标识来生成新种子,这样每个进程的新种子必然不同,后续随机序列也就独立了。这里有两个好用的标识:
方案1:用进程ID(PID)生成新种子
这个方法直接利用操作系统给每个进程分配的唯一PID,结合原种子生成新种子,代码修改如下:
import numpy as np import multiprocessing import os # 用来获取进程ID class mp_worker_class(): def __init__(self,): pass @classmethod def start(self, nb=None, seed=None, nbcore=None): lfp_p=np.empty((nbcore,nb)) pipe_list = [] for h in range(nbcore): recv_end, send_end = multiprocessing.Pipe( ) p = multiprocessing.Process(target=self.mp_worker , args=(h, nb, seed, send_end )) p.start() pipe_list.append(recv_end) for idx, recv_end in enumerate(pipe_list): lfp_p[idx,:]=recv_end.recv() return lfp_p @classmethod def mp_worker(self,h, nb=None, seed=None, send_end=None): np.random.seed(seed) np.random.seed(0) # 统一固定种子,生成第一组一致的随机数 print(h,np.random.rand(5)) # 用进程ID+原种子生成唯一新种子 process_id = os.getpid() # 哈希后取模,确保在numpy种子的有效范围(0到2^32-1)内 new_seed = hash((seed, process_id)) % (2**32 - 1) np.random.seed(new_seed) print(h, np.random.rand(5)) send_end.send(np.random.rand(5)) return if __name__ == '__main__': print(mp_worker_class().start(nb=10, seed=1, nbcore=3 ))
方案2:用进程索引生成新种子(更适配jitclass)
如果@jitclass限制你不能调用os.getpid(),那直接用创建进程时的索引h就行——因为每个进程的h是唯一的(0、1、2...),结合原种子生成新种子,完全不需要额外模块:
# 在mp_worker方法里替换种子生成部分 new_seed = hash((seed, h)) % (2**32 - 1) np.random.seed(new_seed)
为什么这方法管用?
不管是PID还是进程索引,都是每个进程独有的值,所以生成的new_seed在每个进程里都不一样,调用np.random.seed(new_seed)后,每个进程的随机数生成器状态就彻底独立了。而且这种方式完全不依赖时间或者np.random.seed(None),完美适配@jitclass的限制。
修改后的预期输出
你会看到第一组随机数还是完全一致,第二组每个进程都不一样了,比如类似这样:
0 [0.5488135 0.71518937 0.60276338 0.54488318 0.4236548 ] 0 [0.37591933 0.43441666 0.81303047 0.00949575 0.40894335] 1 [0.5488135 0.71518937 0.60276338 0.54488318 0.4236548 ] 1 [0.93482738 0.86919454 0.83116105 0.06308839 0.20434528] 2 [0.5488135 0.71518937 0.60276338 0.54488318 0.4236548 ] 2 [0.62740647 0.45683322 0.82071748 0.08931912 0.75944543]
内容的提问来源于stack exchange,提问作者ymmx

