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

mpi4py并行代码中while循环满足条件却无法终止的问题

MPI4py并行骰子模拟程序挂起问题修复

我正在用mpi4py进行并行化练习,实现投掷2个骰子指定次数(按进程数拆分,即npp)并统计点数的功能,结果存入字典,计算均值偏差,直到mean_dev小于0.001时终止程序。所有功能逻辑正常,但代码在满足终止条件后无法退出,出现挂起现象。

原问题代码

from ctypes.wintypes import SIZE
from dice import * # 生成字典的自定义类
from random import randint
import matplotlib.pyplot as plt
import numpy as np
from mpi4py import MPI
from math import sqrt

def simulation(f_events, f_sides, f_n_dice):
    f_X = dice(sides, n_dice).myDice() # 嵌套字典,最后一层存储点数和的统计
    for j in range(f_events): # 处理所有投掷事件
        n = [] # 存储每个骰子的点数
        for i in range(1, f_n_dice+1): # 遍历每个骰子
            k = randint(1, f_sides) # 生成随机点数
            n.append(k)
            f_X[i][k] += 1 # 对应骰子的点数计数+1
        sum_throw = sum(n) # 本次投掷的点数和
        f_X[f_n_dice+1][sum_throw] += 1 # 点数和的计数+1
    return f_X

npp = int(4)//4 # 每个进程处理的投掷次数
sides = 6 # 骰子面数
n_dice = 2 # 骰子数量

comm = MPI.COMM_WORLD # 通信器
rank = comm.Get_rank() # 当前进程编号
size = comm.Get_size() # 总进程数

# -------------------- 并行化核心代码 --------------------#

seq = (2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12)
AUX = dict.fromkeys(seq, 0)
mean_dev = 1
while True:
    msg = comm.bcast(npp, root = 0)
    print("---> msg: ", msg, " for rank ", rank)
    print("The mean dev for %d" %rank + " is: ", mean_dev)

    D = simulation(npp, sides, n_dice)
    
    Dp = comm.gather(D, root = 0)
    print("This is Dp: ", Dp)
    
    summ = 0
    prob = [1/36, 2/36, 3/36, 4/36, 5/36, 6/36, 5/36, 4/36, 3/36, 2/36, 1/36]

    if rank==0:
        for p in range(0, size): 
                for n in range(dice().min, dice().max+1): # 遍历点数和的取值范围
                    AUX[n] += Dp[p][n_dice+1][n] # 累加所有进程的统计结果
                print(Dp[p][n_dice+1])
    
        print("The final dictionary is: ", AUX, sum(AUX[j] for j in AUX))

        for i in range(dice().min, dice().max+1):
            exp = (prob[i-2])*(sum(AUX[j] for j in AUX))
            x = (AUX[i]-exp)/exp
            summ = summ + pow(x, 2)

        mean_dev = (1/11)*sqrt(summ)
        print("The deviation for {} is {}.".format(sum(AUX[j] for j in AUX), mean_dev))

    if mean_dev > 0.001:
        npp = 2*npp
        # new_msg = comm.bcast(npp, root = 0)
        # print("---> new_msg: ", new_msg, " for rank ", rank)
    else:
        break

修改后的代码(经@victor-eijkhout建议)

from ctypes.wintypes import SIZE
from dice import *
from random import randint
import matplotlib.pyplot as plt
import numpy as np
from mpi4py import MPI
from math import sqrt

def simulation(f_events, f_sides, f_n_dice):
    f_X = dice(sides, n_dice).myDice() # 嵌套字典,最后一层存储点数和的统计
    for j in range(f_events): # 处理所有投掷事件
        n = [] # 存储每个骰子的点数
        for i in range(1, f_n_dice+1): # 遍历每个骰子
            k = randint(1, f_sides) # 生成随机点数
            n.append(k)
            f_X[i][k] += 1 # 对应骰子的点数计数+1
        sum_throw = sum(n) # 本次投掷的点数和
        f_X[f_n_dice+1][sum_throw] += 1 # 点数和的计数+1
    return f_X

npp = int(4)//4 # 每个进程处理的投掷次数
sides = 6 # 骰子面数
n_dice = 2 # 骰子数量

comm = MPI.COMM_WORLD # 通信器
rank = comm.Get_rank() # 当前进程编号
size = comm.Get_size() # 总进程数

# -------------------- 并行化核心代码 --------------------#

seq = (2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12)
AUX = dict.fromkeys(seq, 0)
mean_dev = 1
while True:
    msg = comm.bcast(npp, root = 0)
    #print("---> msg: ", msg, " for rank ", rank)
    
    D = simulation(npp, sides, n_dice)
        
    Dp = comm.gather(D, root = 0)
    #if Dp != None: print("This is Dp: ", Dp)

    
    #print("The mean dev for %d" %rank + " is: ", mean_dev)

    if rank==0:
        
        summ = 0
        prob = [1/36, 2/36, 3/36, 4/36, 5/36, 6/36, 5/36, 4/36, 3/36, 2/36, 1/36]

        for p in range(0, size): 
                for n in range(dice().min, dice().max+1): # 遍历点数和的取值范围
                    AUX[n] += Dp[p][n_dice+1][n] # 累加所有进程的统计结果
                print(Dp[p][n_dice+1])
    
        print("The final dictionary is: ", AUX, sum(AUX[j] for j in AUX))

        for i in range(dice().min, dice().max+1):
            exp = (prob[i-2])*(sum(AUX[j] for j in AUX))
            x = (AUX[i]-exp)/exp
            summ = summ + pow(x, 2)

        mean_dev = (1/11)*sqrt(summ)
        print("The deviation for {} is {}.".format(sum(AUX[j] for j in AUX), mean_dev))

    #new_mean_dev = comm.gather(mean_dev, root = 0)
    new_mean_dev = comm.bcast(mean_dev, root = 0)
    print("---> msg2: ", new_mean_dev, " for rank ", rank)

    if new_mean_dev < 0.001:
        break
        # new_msg = comm.bcast(npp, root = 0)
        # print("---&gt; new_msg: ", new_msg, " for rank ", rank)
        
    else:
        npp = 2*npp
        print("The new npp is: ", npp)

内容的提问来源于stack exchange,提问作者Carreira

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 21:01:20