Skip to content

One code path for serial and MPI runs

Download notebookView on GitHubRun when the site was built.

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.

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: False
serial stand-in: True

Under 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.2
host runnervm8df0l
launched under MPI False
launcher variable -
MAYBEMPI override -
local rank 0
mpi4py installed False
MPI serial stand-in
rank 0 of 1

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_WORLD
rng = 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.9520

Buffer 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 process
print(histogram)
[19. 19. 26. 14. 22.]
[19. 19. 26. 14. 22.]

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])

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 mpi4py
print(serial.COMM_WORLD.allgather("only me"))
['only me']