在似然采样器内首次调用pool.starmap()出现挂起问题如何解决
问题根因
- 进程池对象无法被序列化,不能作为参数传递给likelihood函数:绝大多数MC采样器(如emcee等)会将传入的似然函数及关联参数做pickle序列化后分发调用,而进程池包含内核级句柄资源,无法被序列化,子进程拿到损坏的对象后直接进入挂起状态,无CPU占用。
- 进程池初始化位置不符合spawn启动模式要求:Windows、macOS系统默认的multiprocessing启动模式为spawn,子进程启动时会重新导入主模块,全局位置初始化的进程池会导致子进程递归创建新的进程池,触发死锁。
- 大量复杂对象重复序列化开销:每次调用
starmap时都需要将42个序列化后的复杂函数对象传给子进程,若对象包含锁、文件句柄等不可序列化属性,会直接导致子进程挂起。
修复方案
- 将进程池、list_objs等资源的初始化放到
if __name__ == "__main__"块内,避免子进程重复创建 - 用偏函数将进程池、list_objs绑定到似然函数的闭包内,不要作为参数直接传给采样器,避免被序列化
- 可选优化:用进程池的
initializer将list_objs一次性加载到每个子进程的全局内存,不需要每次调用都重复序列化传递,大幅降低开销
基础修复版代码
import multiprocessing import numpy as np import itertools from functools import partial def pred_per_obj(obj, x1): return obj.predict(x1) def likelihood(x1, mypool, list_objs): xi_pred = mypool.starmap(pred_per_obj, zip(list_objs, itertools.repeat(x1))) return np.mean(xi_pred) if __name__ == "__main__": # 所有全局资源都放在主进程块内初始化 list_objs = 此处替换为加载42个序列化对象的代码 mypool = multiprocessing.Pool(4) # 用偏函数绑定参数,不需要把进程池传给采样器 bounded_likelihood = partial(likelihood, mypool=mypool, list_objs=list_objs) MCsampler.sample(bounded_likelihood, priors) # 采样结束后关闭进程池 mypool.close() mypool.join()
更高性能优化版代码
如果对象体积大、调用频次高,可提前把对象加载到子进程内存,避免重复序列化开销:
import multiprocessing import numpy as np import itertools from functools import partial # 子进程全局变量存储预加载的对象列表 worker_list_objs = None def init_worker(list_objs): global worker_list_objs worker_list_objs = list_objs def pred_per_obj(x1, obj_idx): return worker_list_objs[obj_idx].predict(x1) def likelihood(x1, mypool, obj_num): xi_pred = mypool.starmap(pred_per_obj, zip(itertools.repeat(x1), range(obj_num))) return np.mean(xi_pred) if __name__ == "__main__": list_objs = 此处替换为加载42个序列化对象的代码 # 初始化时把对象列表一次性传给所有子进程 mypool = multiprocessing.Pool(4, initializer=init_worker, initargs=(list_objs,)) bounded_likelihood = partial(likelihood, mypool=mypool, obj_num=len(list_objs)) MCsampler.sample(bounded_likelihood, priors) mypool.close() mypool.join()
内容的提问来源于stack exchange,提问作者Sandy Yuan
相关产品推荐
相关产品推荐

