mpi4py结合SciPy minimize迭代后仅rank0参与,如何重激活其他进程?
问题描述
使用SciPy的minimize结合MPI并行计算时,第一次迭代后仅rank0进程参与目标函数(Objective)的评估,非0rank的工作进程完成初始任务后不再被调用,如何让这些工作进程在后续迭代中继续参与计算?
原始代码
from scipy.optimize import minimize, OptimizeResult from mpi4py import MPI import numpy as np import logging logging.basicConfig(filename='job.log', level=logging.INFO) class Solver(): def __init__(self, SampleTimes, InitialArray): self.comm = MPI.COMM_WORLD self.rank = self.comm.Get_rank() self.size = self.comm.Get_size() self.SampleTimes = SampleTimes self.InitialArray = InitialArray self.Max = None def f(self, x_): return np.sum(x_) def Objective(self, x): logging.info(f"Entering Objective on rank {self.rank}") self.x = self.comm.bcast(x if self.rank == 0 else None, root=0) logging.info(f"Logging x: {self.x}") if not isinstance(self.x, OptimizeResult): tstep_select = np.array_split(self.SampleTimes, self.size)[self.rank] local_results = [] for t in tstep_select: logging.info(f"Processing t={t} on rank {self.rank}") result = t*self.f(self.x) local_results.append( (t, result) ) logging.info(f"Response for t={t}: {local_results[-1][-1]}") all_results = self.comm.gather(local_results, root=0) if self.rank==0: all_results = [item for sublist in all_results for item in sublist] all_results = np.array(all_results) all_results = all_results[all_results[:,0].argsort()] scalar = np.trapz(all_results[:,1], all_results[:,0]) return -scalar def Maximize(self,): if self.rank == 0: self.Max = minimize(self.Objective, self.InitialArray) print(self.Max) else: while not self.Max: self.Objective(None) self.Max = self.comm.bcast(self.Max if self.rank==0 else None, root=0) if __name__=='__main__': t_eval = np.linspace(0, 100, 100) a_init = np.random.rand(10) Instance = Solver(SampleTimes=t_eval, InitialArray=a_init) Instance.Maximize()
问题根源
非0rank进程的while not self.Max循环中,反复调用Objective(None),但此时rank0正在minimize内部调用Objective传递新的x值,两者的MPI通信逻辑不匹配:非0rank在自己的Objective(None)调用中执行bcast,而rank0此时的Objective调用也在执行bcast,导致通信冲突或同步失败,后续迭代中非0rank无法正确接收rank0广播的新x值,也就无法参与计算。
解决方案
重构代码,让所有进程在每次迭代时同步协作:rank0负责驱动优化,广播新的x值;非0rank持续等待rank0的广播指令,要么执行计算任务,要么接收结束信号。
修改后的代码
from scipy.optimize import minimize, OptimizeResult from mpi4py import MPI import numpy as np import logging logging.basicConfig(filename='job.log', level=logging.INFO) class Solver(): def __init__(self, SampleTimes, InitialArray): self.comm = MPI.COMM_WORLD self.rank = self.comm.Get_rank() self.size = self.comm.Get_size() self.SampleTimes = SampleTimes self.InitialArray = InitialArray self.Max = None # 结束标志,用于通知工作进程停止 self.stop_flag = False def f(self, x_): return np.sum(x_) def Objective(self, x): # 所有进程同步接收rank0广播的x或结束信号 broadcast_data = self.comm.bcast((x, self.stop_flag) if self.rank == 0 else None, root=0) self.x, self.stop_flag = broadcast_data logging.info(f"Rank {self.rank} received x: {self.x}, stop_flag: {self.stop_flag}") # 如果收到结束信号,直接返回 if self.stop_flag: return if not isinstance(self.x, OptimizeResult): tstep_select = np.array_split(self.SampleTimes, self.size)[self.rank] local_results = [] for t in tstep_select: logging.info(f"Processing t={t} on rank {self.rank}") result = t*self.f(self.x) local_results.append( (t, result) ) logging.info(f"Response for t={t}: {local_results[-1][-1]}") all_results = self.comm.gather(local_results, root=0) if self.rank == 0: all_results = [item for sublist in all_results for item in sublist] all_results = np.array(all_results) all_results = all_results[all_results[:,0].argsort()] scalar = np.trapz(all_results[:,1], all_results[:,0]) return -scalar def Maximize(self): if self.rank == 0: # rank0驱动优化 self.Max = minimize(self.Objective, self.InitialArray) print(self.Max) # 优化完成后设置结束标志,广播给所有进程 self.stop_flag = True # 最后一次广播,通知工作进程停止 self.Objective(None) else: # 工作进程持续执行Objective,直到收到结束信号 while not self.stop_flag: self.Objective(None) # 同步优化结果到所有进程 self.Max = self.comm.bcast(self.Max if self.rank == 0 else None, root=0) if __name__=='__main__': t_eval = np.linspace(0, 100, 100) a_init = np.random.rand(10) Instance = Solver(SampleTimes=t_eval, InitialArray=a_init) Instance.Maximize()
关键改动说明
- 同步通信逻辑:在
Objective中,所有进程通过bcast同步接收rank0发送的(x, stop_flag)数据,确保每次迭代时所有进程都能获取到最新的x值或结束信号。 - 结束信号机制:添加
stop_flag变量,当rank0完成优化后,设置该标志并广播,非0rank进程检测到标志后退出循环。 - 工作进程循环:非0rank进程不再依赖
self.Max判断是否停止,而是通过stop_flag同步结束,避免通信不匹配。
内容的提问来源于stack exchange,提问作者Sterling Butters
相关产品推荐
相关产品推荐

