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

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的线程模型冲突:

  1. fork出来的子进程会完全复制父进程的内存空间,包括MPI的通信上下文和文件描述符,但MPI库从设计上就不支持这种操作——它认为每个进程的MPI上下文是唯一的,fork会导致上下文混乱。
  2. 你设置了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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 09:38:12