One code path for serial and MPI runs
A code that runs under mpirun on a cluster often also runs as a plain python script.py on a laptop.
Importing mpi4py.MPI in the serial run starts MPI anyway, which costs close to a second and makes
every collective cost something. This tutorial writes one function that works in both cases.
Get the MPI module
Section titled “Get the MPI module”maybempi.get_mpi() looks at the environment the launcher sets up, without importing mpi4py. This
notebook was not started by mpirun, so it gets the serial stand-in:
import maybempi
MPI = maybempi.get_mpi()print(MPI)print("launched under MPI:", maybempi.launched_under_mpi())print("serial stand-in:", maybempi.is_serial(MPI))<SerialMPI: serial stand-in for mpi4py.MPI>launched under MPI: Falseserial stand-in: TrueUnder mpirun -n 4 python script.py, the same line returns mpi4py.MPI. from maybempi import MPI does the same in one line. The decision is made
once per process. The maybempi command prints what it would decide, and why:
from maybempi.cli import report
print(report(init=True))item value------------------ ---------------maybempi 0.1.2host runnervm8df0llaunched under MPI Falselauncher variable -MAYBEMPI override -local rank 0mpi4py installed FalseMPI serial stand-inrank 0 of 1Write the code once
Section titled “Write the code once”The stand-in has the module attributes and communicator methods a serial run needs, with the
results MPI gives on one process: allreduce(x) returns x, gather(x) returns [x].
import numpy as np
def global_mean(local_values, comm): """Mean over the values of all ranks.""" total = comm.allreduce(float(np.sum(local_values)), op=MPI.SUM) count = comm.allreduce(len(local_values), op=MPI.SUM) return total / count
comm = MPI.COMM_WORLDrng = np.random.default_rng(seed=comm.Get_rank())values = rng.normal(loc=1.0, size=1000)print( f"rank {comm.Get_rank()} of {comm.Get_size()}: mean = {global_mean(values, comm):.4f}")rank 0 of 1: mean = 0.9520Buffer collectives copy what MPI would copy on one process. With IN_PLACE they do nothing:
histogram = np.bincount(rng.integers(0, 5, size=100), minlength=5).astype(float)global_histogram = np.zeros_like(histogram)comm.Allreduce(histogram, global_histogram, op=MPI.SUM)print(global_histogram)
comm.Allreduce(MPI.IN_PLACE, histogram, op=MPI.SUM) # unchanged on one processprint(histogram)[19. 19. 26. 14. 22.][19. 19. 26. 14. 22.]Missing methods raise
Section titled “Missing methods raise”A stand-in that returns None for everything hides bugs: a halo exchange that silently does
nothing gives wrong results instead of an error. Methods the stand-in does not implement raise
AttributeError:
comm.Create_cart(dims=[1], periods=[True])SerialCartcomm(COMM_WORLD, dims=[1], periods=[1])Force the decision
Section titled “Force the decision”MAYBEMPI=0 keeps a launched process serial, and MAYBEMPI=1 uses MPI under a launcher that
maybempi does not recognize. get_mpi(False) and get_mpi(True) decide in code:
serial = maybempi.get_mpi(False) # never imports mpi4pyprint(serial.COMM_WORLD.allgather("only me"))['only me']