Wrap Open MPI mpi_f08 for Python¶
PRIK reads Open MPI's Fortran mpi_f08 sources and generates a Python
extension that calls the installed Open MPI library, with no hand-written C,
Cython, or ctypes. You edit the generated .pyi contract into the Python
signatures you want and add a short mpi4py-style
module. The result runs this program, which mpi4py also runs unchanged apart
from its import:
import numpy as np
import prik_mpi as MPI
comm = MPI.COMM_WORLD
rank = comm.Get_rank()
size = comm.Get_size()
ROOT = np.int32(0)
TAG = np.int32(77)
# Point to point: rank 0 sends four integers to rank 1.
if rank == 0:
data = np.arange(4, dtype=np.int32)
comm.Send(data, dest=rank + 1, tag=TAG)
elif rank == 1:
data = np.empty(4, dtype=np.int32)
comm.Recv(data, source=rank - 1, tag=TAG)
print(f"rank 1 received {data.tolist()}")
# Broadcast: rank 0's values reach every rank.
data = np.arange(3, dtype=np.int32) if rank == 0 else np.empty(3, dtype=np.int32)
comm.Bcast(data, root=ROOT)
# Reductions: every rank contributes.
values = np.array([rank + 1, rank + 2], dtype=np.int32)
total = np.empty_like(values)
comm.Allreduce(values, total, op=MPI.SUM)
largest = np.empty_like(values)
comm.Reduce(values, largest, op=MPI.MAX, root=ROOT)
comm.Allreduce(MPI.IN_PLACE, values, op=MPI.SUM)
comm.Barrier()
print(f"rank {rank} of {size}: bcast {data.tolist()}, sum {total.tolist()}, in place {values.tolist()}")
if rank == 0:
print(f"rank 0 max {largest.tolist()}")
You build two APIs: the wrapped API (prik_openmpi_f08), generated by
PRIK, and the Python API (prik_mpi.py), a few lines of Python on top of
it. The Python API is deliberately minimal: it sends only MPI_INT buffers and
reports no receive status, to show the pattern rather than all of mpi4py. In
the measured setup, both had lower median times than mpi4py for small
Allreduce calls; see Compare call times for the other
operations and the test environment.
How this works¶
- Select the 18
mpi_f08names the program needs, out of hundreds (step 1). - Generate, then edit, a
.pyicontract. PRIK writes the Fortran signatures as they are; the edits make them Pythonic:ierrorbecomes an exception,countcomes from the NumPy buffer, and output arguments become return values (steps 2-4). - Add a thin Python layer,
prik_mpi.py, that gives the familiarMPI.COMM_WORLDandcomm.Send(...)style (step 5).
What you need¶
- Linux or macOS, Python 3.10 or later, NumPy, and PRIK; see Installation.
- GNU Fortran and GCC of the same version; CI uses version 13.
- Open MPI's source and build trees, configured with the
mpi_f08bindings. PRIK reads thempi_f08Fortran sources, which an installed Open MPI does not ship. If you do not have them, follow Build Open MPI from source first; it also builds a matching mpi4py for the comparison. - The tutorial's files, from the PRIK repository:
mpi_exports.txt,mpi_f08.pyi,prik_mpi.py,mpi_example.py, andopenmpi_f08.py. Download them into an empty working directory:
PRIK_RAW=https://raw.githubusercontent.com/PyNumLab/prik/main
for file in \
tests/fortran/assumed_types/end_to_end/fixtures/contracts/openmpi/mpi_exports.txt \
tests/fortran/assumed_types/end_to_end/fixtures/contracts/openmpi/mpi_f08.pyi \
tests/fortran/assumed_types/end_to_end/fixtures/runtime/prik_mpi.py \
tests/fortran/assumed_types/end_to_end/fixtures/runtime/mpi_example.py \
benchmarks/openmpi_f08.py; do
curl -fsSLO "$PRIK_RAW/$file"
done
Point two variables at Open MPI's trees, and put that Open MPI's mpifort and
mpirun on PATH:
OMPI_SRC="$HOME/openmpi-5.0.11/source"
OMPI_BUILD="$HOME/openmpi-5.0.11/build"
1. Choose what to wrap¶
mpi_exports.txt selects the names the program needs from mpi_f08's
hundreds:
mpi_f08::MPI_Init
mpi_f08::MPI_Finalize
mpi_f08::MPI_Comm_rank
mpi_f08::MPI_Comm_size
mpi_f08::MPI_Barrier
mpi_f08::MPI_Send
mpi_f08::MPI_Recv
mpi_f08::MPI_Bcast
mpi_f08::MPI_Reduce
mpi_f08::MPI_Allreduce
mpi_f08::MPI_COMM_WORLD
mpi_f08::MPI_INT
mpi_f08::MPI_SUM
mpi_f08::MPI_MAX
mpi_f08::MPI_IN_PLACE
mpi_f08::MPI_STATUS_IGNORE
mpi_f08::MPI_ANY_SOURCE
mpi_f08::MPI_ANY_TAG
2. Generate the contract¶
PRIK starts from the entry source and follows each use through both Open MPI
trees. The -I directories are where those sources find their headers:
python3 -m prik generate --pyi \
"$OMPI_SRC/ompi/mpi/fortran/use-mpi-f08/mpi-f08.F90" \
--module-source-dir "$OMPI_SRC" \
--module-source-dir "$OMPI_BUILD" \
--export-symbols mpi_exports.txt \
--out contract \
--compiler mpifort \
-I "$OMPI_BUILD" \
-I "$OMPI_BUILD/ompi/mpi/fortran/use-mpi-f08" \
-I "$OMPI_BUILD/ompi/mpi/fortran/use-mpi-f08/mod" \
-I "$OMPI_SRC" \
-I "$OMPI_BUILD/ompi/include" \
-I "$OMPI_SRC/ompi/include"
contract/ now holds one .pyi file per Fortran module, with Open MPI's
handle types, constants, and routines as declared in Fortran. For example,
contract/mpi_f08_interfaces.pyi has:
@overload("mpi_allreduce_f08")
def mpi_allreduce(
sendbuf: Annotated[AnyNative[Flat], ReadOnly],
recvbuf: AnyNative[Flat],
count: Int32,
datatype: mpi_types.Mpi_Datatype,
op: mpi_types.Mpi_Op,
comm: mpi_types.Mpi_Comm,
ierror: Int32[()] = ...
) -> Returns["ierror", Int32[()]] | None: ...
3. Edit the contract¶
Replace the published module with the edited one you downloaded:
cp mpi_f08.pyi contract/mpi_f08.pyi
The downloaded mpi_f08.pyi is the generated contract after these edits.
Compare allreduce with the generated version in step 2:
@raises(status="ierror", success=0)
@bind("MPI_Allreduce")
@native_call([Arg(0), Arg(1), Int32(Arg(1).size), Arg(2), Arg(3), Arg(4), Hidden("ierror", Int32)])
def allreduce(
sendbuf: Annotated[AnyNative[Flat], ReadOnly],
recvbuf: AnyNative[Flat],
datatype: Mpi_Datatype,
op: Mpi_Op,
comm: Mpi_Comm,
) -> None: ...
Each function still calls the Fortran routine it names; the edits change only its Python signature:
| Edit | Effect in Python |
|---|---|
@bind("MPI_Send") on def send |
The Python name differs from the Fortran name. |
Int32(Arg(0).size) |
The count comes from the buffer, so the caller does not pass it. |
Return("rank", 0) |
The output argument becomes the return value. |
Hidden("ierror", Int32) with @raises(status="ierror", success=0) |
A nonzero error code raises an exception. |
The complete edited mpi_f08.pyi
from prik.contracts import Annotated, AnyNative, Arg, Flat, Hidden, Int32, ReadOnly, Return, bind, native_call, raises
from .mpi_f08_types import (
Mpi_Comm,
Mpi_Datatype,
Mpi_Op,
Mpi_Status,
mpi_any_source,
mpi_any_tag,
mpi_comm_world,
mpi_in_place,
mpi_int,
mpi_max,
mpi_status_ignore,
mpi_sum,
)
@raises(status="ierror", success=0)
@bind("MPI_Init")
@native_call([Hidden("ierror", Int32)])
def init() -> None: ...
@raises(status="ierror", success=0)
@bind("MPI_Finalize")
@native_call([Hidden("ierror", Int32)])
def finalize() -> None: ...
@raises(status="ierror", success=0)
@bind("MPI_Comm_rank")
@native_call([Arg(0), Return("rank", 0), Hidden("ierror", Int32)])
def comm_rank(comm: Mpi_Comm) -> Int32: ...
@raises(status="ierror", success=0)
@bind("MPI_Comm_size")
@native_call([Arg(0), Return("size", 0), Hidden("ierror", Int32)])
def comm_size(comm: Mpi_Comm) -> Int32: ...
@raises(status="ierror", success=0)
@bind("MPI_Barrier")
@native_call([Arg(0), Hidden("ierror", Int32)])
def barrier(comm: Mpi_Comm) -> None: ...
@raises(status="ierror", success=0)
@bind("MPI_Send")
@native_call([Arg(0), Int32(Arg(0).size), Arg(1), Arg(2), Arg(3), Arg(4), Hidden("ierror", Int32)])
def send(
buf: Annotated[AnyNative[Flat], ReadOnly],
datatype: Mpi_Datatype,
dest: Int32,
tag: Int32,
comm: Mpi_Comm,
) -> None: ...
@raises(status="ierror", success=0)
@bind("MPI_Recv")
@native_call([Arg(0), Int32(Arg(0).size), Arg(1), Arg(2), Arg(3), Arg(4), Arg(5), Hidden("ierror", Int32)])
def recv(
buf: AnyNative[Flat],
datatype: Mpi_Datatype,
source: Int32,
tag: Int32,
comm: Mpi_Comm,
status: Mpi_Status,
) -> None: ...
@raises(status="ierror", success=0)
@bind("MPI_Bcast")
@native_call([Arg(0), Int32(Arg(0).size), Arg(1), Arg(2), Arg(3), Hidden("ierror", Int32)])
def bcast(buffer: AnyNative[Flat], datatype: Mpi_Datatype, root: Int32, comm: Mpi_Comm) -> None: ...
@raises(status="ierror", success=0)
@bind("MPI_Reduce")
@native_call([Arg(0), Arg(1), Int32(Arg(1).size), Arg(2), Arg(3), Arg(4), Arg(5), Hidden("ierror", Int32)])
def reduce(
sendbuf: Annotated[AnyNative[Flat], ReadOnly],
recvbuf: AnyNative[Flat],
datatype: Mpi_Datatype,
op: Mpi_Op,
root: Int32,
comm: Mpi_Comm,
) -> None: ...
@raises(status="ierror", success=0)
@bind("MPI_Allreduce")
@native_call([Arg(0), Arg(1), Int32(Arg(1).size), Arg(2), Arg(3), Arg(4), Hidden("ierror", Int32)])
def allreduce(
sendbuf: Annotated[AnyNative[Flat], ReadOnly],
recvbuf: AnyNative[Flat],
datatype: Mpi_Datatype,
op: Mpi_Op,
comm: Mpi_Comm,
) -> None: ...
__all__ = [
"init",
"finalize",
"comm_rank",
"comm_size",
"barrier",
"send",
"recv",
"bcast",
"reduce",
"allreduce",
"Mpi_Comm",
"Mpi_Datatype",
"Mpi_Op",
"Mpi_Status",
"mpi_any_source",
"mpi_any_tag",
"mpi_comm_world",
"mpi_in_place",
"mpi_int",
"mpi_max",
"mpi_status_ignore",
"mpi_sum",
]
4. Build the wrapped API¶
python3 -m prik contract/__init__.pyi \
--compiler "$(mpifort --showme:command)" \
--wrapper-fortran-flags="$(mpifort --showme:compile)" \
--native-library $(mpifort --showme:libs) \
--native-library-dir $(mpifort --showme:libdirs) \
--out prik_openmpi_f08 \
--out-dir build
This writes prik_openmpi_f08.so in the working directory. Only PRIK's
generated code is compiled; the extension links to the installed Open MPI.
5. Add the Python API¶
prik_mpi.py imports the wrapped API as _mpi and gives it mpi4py's Comm
object. Each method forwards to one wrapped function, adding the datatype and
the communicator handle:
class Comm:
"""A communicator, with the methods mpi4py spells for it."""
def __init__(self, handle):
self.handle = handle
def Get_rank(self):
return _mpi.comm_rank(self.handle)
def Allreduce(self, sendbuf, recvbuf, op=SUM):
_mpi.allreduce(sendbuf, recvbuf, _INT, op, self.handle)
COMM_WORLD = Comm(_mpi.mpi_comm_world)
# Like mpi4py, MPI starts when this module is imported and stops at exit.
_mpi.init()
atexit.register(_mpi.finalize)
The complete prik_mpi.py
"""An mpi4py-style Python API over the PRIK-generated Open MPI extension.
It illustrates the shape of mpi4py rather than all of it: every buffer is an
np.int32 array sent as MPI_INT, and no receive reports a status.
"""
import atexit
import numpy as np
from prik_openmpi_f08 import mpi_f08 as _mpi
ANY_SOURCE = _mpi.mpi_any_source
ANY_TAG = _mpi.mpi_any_tag
IN_PLACE = _mpi.mpi_in_place
SUM = _mpi.mpi_sum
MAX = _mpi.mpi_max
_ZERO = np.int32(0)
_INT = _mpi.mpi_int
_STATUS_IGNORE = _mpi.mpi_status_ignore
class Comm:
"""A communicator, with the methods mpi4py spells for it."""
def __init__(self, handle):
self.handle = handle
def Get_rank(self):
return _mpi.comm_rank(self.handle)
def Get_size(self):
return _mpi.comm_size(self.handle)
def Barrier(self):
_mpi.barrier(self.handle)
def Send(self, buf, dest, tag=_ZERO):
_mpi.send(buf, _INT, dest, tag, self.handle)
def Recv(self, buf, source=ANY_SOURCE, tag=ANY_TAG):
_mpi.recv(buf, _INT, source, tag, self.handle, _STATUS_IGNORE)
def Bcast(self, buf, root=_ZERO):
_mpi.bcast(buf, _INT, root, self.handle)
def Reduce(self, sendbuf, recvbuf, op=SUM, root=_ZERO):
_mpi.reduce(sendbuf, recvbuf, _INT, op, root, self.handle)
def Allreduce(self, sendbuf, recvbuf, op=SUM):
_mpi.allreduce(sendbuf, recvbuf, _INT, op, self.handle)
COMM_WORLD = Comm(_mpi.mpi_comm_world)
# Like mpi4py, MPI starts when this module is imported and stops at exit.
_mpi.init()
atexit.register(_mpi.finalize)
6. Run it¶
mpirun -n 2 python3 mpi_example.py
The ranks print, in either order:
rank 1 received [0, 1, 2, 3]
rank 0 of 2: bcast [0, 1, 2], sum [3, 5], in place [3, 5]
rank 0 max [2, 3]
rank 1 of 2: bcast [0, 1, 2], sum [3, 5], in place [3, 5]
The same program under mpi4py prints the same lines:
sed 's/^import prik_mpi as MPI$/from mpi4py import MPI/' mpi_example.py > mpi4py_example.py
mpirun -n 2 python3 mpi4py_example.py
Compare call times¶
for api in wrapped python mpi4py; do
mpirun -n 2 python3 openmpi_f08.py "$api"
done
openmpi_f08.py times every API the same way with mpi4py's MPI.Wtime. On an
AMD Ryzen 5 5600H with Ubuntu 22.04.5, Open MPI 5.0.11, and mpi4py 4.1.2, the
median of five runs was:
| Operation | mpi4py | Wrapped API | Python API |
|---|---|---|---|
Allreduce, 1 int32 |
1.037 µs | 0.697 µs (33% faster) | 0.841 µs (19% faster) |
Allreduce, 1,024 int32 values |
2.314 µs | 1.779 µs (23% faster) | 2.027 µs (12% faster) |
Allreduce, 1,048,576 int32 values |
2.719 ms | 2.752 ms (about the same) | 2.827 ms (about the same) |
Barrier |
0.330 µs | 0.328 µs (about the same) | 0.428 µs (30% slower) |
Get_rank |
30 ns | 145 ns (about 5× slower) | 236 ns (about 8× slower) |
Get_rank does almost nothing, so its time is the overhead of a call, which
PRIK can still reduce.
Build Open MPI from source¶
Do this once if you do not already have Open MPI's source and build trees. It builds Open MPI 5.0.11 and keeps both trees, which PRIK reads. From your working directory:
OMPI_ROOT="$HOME/openmpi-5.0.11"
TUTORIAL_DIR="$PWD"
mkdir -p "$OMPI_ROOT/source" "$OMPI_ROOT/build" "$OMPI_ROOT/toolchain"
ln -sf "$(command -v gfortran-13)" "$OMPI_ROOT/toolchain/gfortran"
ln -sf "$(command -v gcc-13)" "$OMPI_ROOT/toolchain/gcc"
export PATH="$OMPI_ROOT/toolchain:$PATH"
curl -fsSLO https://download.open-mpi.org/release/open-mpi/v5.0/openmpi-5.0.11.tar.bz2
tar -xjf openmpi-5.0.11.tar.bz2 -C "$OMPI_ROOT/source" --strip-components=1
cd "$OMPI_ROOT/build"
../source/configure --prefix="$OMPI_ROOT/install" --enable-mpi-fortran=usempif08 CC=gcc FC=gfortran
make -j2 && make install
cd "$TUTORIAL_DIR"
export PATH="$OMPI_ROOT/install/bin:$PATH"
export LD_LIBRARY_PATH="$OMPI_ROOT/install/lib${LD_LIBRARY_PATH:+:$LD_LIBRARY_PATH}"
export DYLD_LIBRARY_PATH="$OMPI_ROOT/install/lib${DYLD_LIBRARY_PATH:+:$DYLD_LIBRARY_PATH}"
OMPI_SRC="$OMPI_ROOT/source"
OMPI_BUILD="$OMPI_ROOT/build"
Build mpi4py against the same Open MPI; both checks should report 5.0.11:
MPI4PY_BUILD_MPICC="$OMPI_ROOT/install/bin/mpicc" \
python3 -m pip install --no-cache-dir --no-binary=mpi4py mpi4py==4.1.2
mpirun --version
python3 -c 'from mpi4py import MPI; print(MPI.Get_library_version())'
For Open MPI 4.1.8, which CI also tests, change 5.0.11 to 4.1.8 and v5.0
to v4.1.