Python多进程结合mpi4py跨节点通信失效问题排查
问题:Python multiprocessing + mpi4py 跨节点通信失效
嘿,这个问题我太熟了!之前在把单节点multiprocessing代码改成mpi4py跨节点执行时,踩过一模一样的坑——启动子进程后MPI直接卡壳,完全和你描述的场景一致:2节点各跑1个MPI进程,rank1疯狂发数据,rank0收到第一行就停滞不动。
复现代码
你提供的测试代码我整理了格式(修正了导入语句的语法错误):
import os import psutil import multiprocessing import numpy as np import Queue import time from mpi4py import rc rc.initialize = False # if = True, The Init is done when "from mpi4py import MPI" is called rc.thread_level = 'funneled' from mpi4py import MPI def infiniteloop(arg): while True: print(arg) time.sleep(1) # Check if the worker think it's the thread who called MPI.Init() print("Worker is Main Thread %s" %(MPI.Is_thread_main())) print("Rank %d on %s, Process PID for worker = %d" %(MPI.COMM_WORLD.Get_rank(),MPI.Get_processor_name(),os.getpid())) if __name__ == '__main__': MPI.Init() # In the code I'm working on, MPI.Init() has to be done before the miltiprocess initialization proc = multiprocessing.Process(target=infiniteloop, args=('RunningWorker',)) proc.start() print("MultiProcess Stared") comm = MPI.COMM_WORLD rank = comm.Get_rank() size_mpi = comm.Get_size() while True: print("Running Main Thread") print("Main Thread is Main Thread %s" %( MPI.Is_thread_main())) print("Rank %d on %s, Process PID for main = %s" %(MPI.COMM_WORLD.Get_rank(),MPI.Get_processor_name(),os.getpid())) print("Rank %d on %s, rc.thread_level = %s" %((MPI.COMM_WORLD.Get_rank(),MPI.Get_processor_name(), rc.thread_level))) time.sleep(1) # Start MPI Communication, It is just an example of 2D array communication which I know it works print("Start MPI Transfert") #*************** Multiple SEND AND RECEIVE for 2D Array fill randomly SumPsfmean = None TransPsfmean = None TABSIZE = 100 # Create 100x100 array with random np.float64 values # (Because it's really close from the case I'm intersting for) # Row per row communication if rank == 0: psfmean = np.random.rand(TABSIZE,TABSIZE) print(psfmean.dtype) else: psfmean = np.random.rand(TABSIZE,TABSIZE) psfmean_shape = psfmean.shape if rank == 0: SumPsfmean = np.array(range(psfmean.size*(size_mpi-1)), dtype = np.float64) SumPsfmean.shape = (size_mpi-1, psfmean_shape[0], psfmean_shape[1]) TransPsfmean = np.array(range(psfmean[0].size), dtype = np.float64) for i in range(psfmean_shape[1]): print("Rank %d : Send&Receive nb %d" %(rank, i)) if rank == 0: comm.Recv(TransPsfmean, source=1, tag=i) elif rank ==1: comm.Send(psfmean[i], dest=0, tag=i) print("End Send&Receive %d" %i) if rank == 0: for k in range(size_mpi): if k != 0: SumPsfmean[k-1][i] = TransPsfmean proc.join()
运行结果
你的终端输出整理后如下:
[0] MultiProcess Stared [0] Running Main Thread [0] Main Thread is Main Thread : True [0] Rank 0 on genji271, Process PID for main = 227040 [0] Rank 0 on genji271, rc.thread_level = funneled [1] MultiProcess Stared [1] Running Main Thread [1] Main Thread is Main Thread : True [1] Rank 1 on genji272, Process PID for main = 211028 [1] Rank 1 on genji272, rc.thread_level = funneled [0] RunningWorker[0] [1] RunningWorker [0] Start MPI Transfert [0] float64 [1] Start MPI Transfert [1] Rank 1 : Send&Receive nb 0 [1] End Send&Receive 0 [1] Rank 1 : Send&Receive nb 1 [1] End Send&Receive 1 ... # rank1持续发送直到第15行 [0] Rank 0 : Send&Receive nb 0 [0] Worker is Main Thread : True [0] Rank 0 on genji271, Process PID for worker = 227046 [0] RunningWorker [1] Worker is Main Thread : True [1] Rank 1 on genji272, Process PID for worker = 211034 [1] RunningWorker
从结果能明显看到:rank1一直在发送数据,但rank0刚启动接收第0行就停滞了,通信彻底失效。
问题根源分析
核心问题出在Python multiprocessing的默认启动模式(Unix下是fork)和mpi4py的线程模型冲突:
fork出来的子进程会完全复制父进程的内存空间,包括MPI的通信上下文和文件描述符,但MPI库从设计上就不支持这种操作——它认为每个进程的MPI上下文是唯一的,fork会导致上下文混乱。- 你设置了
rc.thread_level = 'funneled',这个模式下只有调用MPI.Init()的主线程才能执行MPI操作,但fork出来的子进程会错误地认为自己是“主线程”(MPI.Is_thread_main()返回True),相当于伪造了MPI状态,直接干扰了父进程的通信流程。
解决方案
方案1:先启动子进程,再初始化MPI(推荐)
mpi4py官方明确建议不要在MPI.Init()之后使用fork,所以最稳妥的方式是调整代码顺序:先启动所有子进程,再初始化MPI。这样子进程完全不会接触到MPI上下文,不会干扰主进程的通信:
import os import multiprocessing import numpy as np import time from mpi4py import rc rc.initialize = False rc.thread_level = 'funneled' from mpi4py import MPI def infiniteloop(arg): while True: print(arg) time.sleep(1) # 子进程里不要调用任何MPI相关函数! print("Worker PID = %d" %(os.getpid())) if __name__ == '__main__': # 先启动子进程 proc = multiprocessing.Process(target=infiniteloop, args=('RunningWorker',)) proc.start() print("MultiProcess Started") # 再初始化MPI MPI.Init() comm = MPI.COMM_WORLD rank = comm.Get_rank() size_mpi = comm.Get_size() # 后续MPI通信逻辑 print("Running Main Thread") print("Rank %d on %s, Process PID for main = %s" %(rank, MPI.Get_processor_name(), os.getpid())) print("Start MPI Transfert") TABSIZE = 100 psfmean = np.random.rand(TABSIZE, TABSIZE).astype(np.float64) psfmean_shape = psfmean.shape if rank == 0: SumPsfmean = np.zeros((size_mpi-1, *psfmean_shape), dtype=np.float64) TransPsfmean = np.zeros(psfmean_shape[1], dtype=np.float64) for i in range(psfmean_shape[1]): print(f"Rank {rank} : Send&Receive nb {i}") if rank == 0: comm.Recv(TransPsfmean, source=1, tag=i) SumPsfmean[0][i] = TransPsfmean elif rank == 1: comm.Send(psfmean[i], dest=0, tag=i) print(f"End Send&Receive {i}") proc.join() MPI.Finalize()
方案2:强制multiprocessing使用spawn模式(适合必须先初始化MPI的场景)
如果你的业务逻辑必须先启动MPI再开子进程,可以强制multiprocessing使用spawn模式——这种模式会启动全新的Python进程,不会继承父进程的MPI上下文,从根源避免干扰:
if __name__ == '__main__': # 强制使用spawn模式,必须放在所有multiprocessing操作之前 multiprocessing.set_start_method('spawn') MPI.Init() proc = multiprocessing.Process(target=infiniteloop, args=('RunningWorker',)) proc.start() # ... 后续代码不变
⚠️ 注意:spawn模式下,子进程无法继承父进程的MPI状态,所以子进程里绝对不能调用任何MPI相关函数(比如原代码里的MPI.Is_thread_main()要删掉,否则会报错)。
验证效果
用上面两种方法修改后,rank0和rank1的MPI通信会正常完成,子进程的输出也不会再干扰主进程的MPI上下文,彻底解决停滞问题。
内容的提问来源于stack exchange,提问作者Nicolas Monnier
相关产品推荐
相关产品推荐

