maybempi 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
maybempi/__init__.py ADDED
@@ -0,0 +1,66 @@
1
+ """Use MPI only when the process was launched under MPI, and a serial stand-in otherwise.
2
+
3
+ Importing ``mpi4py.MPI`` starts MPI, which costs close to a second and makes every
4
+ collective cost something, even on one process. maybempi decides from the
5
+ environment the launcher sets up, without importing mpi4py, and returns either
6
+ ``mpi4py.MPI`` or a serial stand-in with the same interface::
7
+
8
+ import maybempi
9
+
10
+ MPI = maybempi.get_mpi() # mpi4py.MPI under mpirun/srun, else the stand-in
11
+ comm = MPI.COMM_WORLD
12
+ total = comm.allreduce(local_total, op=MPI.SUM)
13
+
14
+ ``from maybempi import MPI`` is the drop-in replacement for ``from mpi4py import
15
+ MPI``: the name resolves to ``get_mpi()`` when it is first imported.
16
+
17
+ :mod:`maybempi.launch` holds the detection, :mod:`maybempi.serial` the stand-in.
18
+ """
19
+
20
+ from typing import Any
21
+
22
+ from maybempi.launch import (
23
+ LAUNCHER_VARIABLES,
24
+ LOCAL_RANK_VARIABLES,
25
+ OVERRIDE_VARIABLE,
26
+ get_mpi,
27
+ is_serial,
28
+ launched_under_mpi,
29
+ launcher_variable,
30
+ local_rank,
31
+ )
32
+ from maybempi.serial import (
33
+ SerialComm,
34
+ SerialMPI,
35
+ SerialPrequest,
36
+ SerialRequest,
37
+ SerialStatus,
38
+ set_copy_hook,
39
+ )
40
+
41
+ __version__ = "0.1.0"
42
+
43
+ __all__ = [
44
+ "LAUNCHER_VARIABLES",
45
+ "LOCAL_RANK_VARIABLES",
46
+ "OVERRIDE_VARIABLE",
47
+ "SerialComm",
48
+ "SerialMPI",
49
+ "SerialPrequest",
50
+ "SerialRequest",
51
+ "SerialStatus",
52
+ "__version__",
53
+ "get_mpi",
54
+ "is_serial",
55
+ "launched_under_mpi",
56
+ "launcher_variable",
57
+ "local_rank",
58
+ "set_copy_hook",
59
+ ]
60
+
61
+
62
+ def __getattr__(name: str) -> Any:
63
+ """Resolve ``maybempi.MPI`` (and ``from maybempi import MPI``) to :func:`get_mpi`."""
64
+ if name == "MPI":
65
+ return get_mpi()
66
+ raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
maybempi/__main__.py ADDED
@@ -0,0 +1,5 @@
1
+ """``python -m maybempi``: the same as the ``maybempi`` command."""
2
+
3
+ from maybempi.cli import main
4
+
5
+ main()
maybempi/cli.py ADDED
@@ -0,0 +1,142 @@
1
+ """The ``maybempi`` command: show what maybempi decides in this environment.
2
+
3
+ Run it the way the application is run, to see whether it would use MPI::
4
+
5
+ maybempi # serial: one table
6
+ mpirun -n 2 maybempi # one table per rank (MPI is not started)
7
+ mpirun -n 2 maybempi --init # start MPI: rank 0 prints one table for all ranks
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import argparse
13
+ import importlib.util
14
+ import os
15
+ import socket
16
+ from collections.abc import Sequence
17
+ from typing import Any
18
+
19
+ from maybempi import __version__
20
+ from maybempi.launch import (
21
+ OVERRIDE_VARIABLE,
22
+ get_mpi,
23
+ is_serial,
24
+ launched_under_mpi,
25
+ launcher_variable,
26
+ local_rank,
27
+ )
28
+
29
+ # columns of the per-rank table, in this order; other items become columns only
30
+ # if they differ between ranks
31
+ _RANK_COLUMNS = ("rank", "host", "local rank", "launcher variable")
32
+
33
+
34
+ def info(init: bool = False) -> dict[str, Any]:
35
+ """Collect what maybempi decides in this process, and why.
36
+
37
+ Args:
38
+ init: Also call :func:`~maybempi.launch.get_mpi` (which starts MPI under a
39
+ launcher) and add the rank and size of ``COMM_WORLD``.
40
+
41
+ Returns:
42
+ The items of the report, by name.
43
+ """
44
+ items: dict[str, Any] = {
45
+ "maybempi": __version__,
46
+ "host": socket.gethostname(),
47
+ "launched under MPI": launched_under_mpi(),
48
+ "launcher variable": launcher_variable() or "-",
49
+ f"{OVERRIDE_VARIABLE} override": os.environ.get(OVERRIDE_VARIABLE, "-"),
50
+ "local rank": local_rank(),
51
+ "mpi4py installed": importlib.util.find_spec("mpi4py") is not None,
52
+ }
53
+ if init:
54
+ mpi = get_mpi()
55
+ comm = mpi.COMM_WORLD
56
+ items["MPI"] = "serial stand-in" if is_serial(mpi) else "mpi4py"
57
+ items["rank"] = comm.Get_rank()
58
+ items["size"] = comm.Get_size()
59
+ return items
60
+
61
+
62
+ def _table(header: Sequence[str], rows: Sequence[Sequence[Any]]) -> str:
63
+ """Format `rows` as a plain-text table with aligned columns."""
64
+ cells = [[str(value) for value in row] for row in rows]
65
+ widths = [
66
+ max(len(str(h)), *(len(row[i]) for row in cells)) for i, h in enumerate(header)
67
+ ]
68
+
69
+ def line(values: Sequence[str]) -> str:
70
+ return " ".join(
71
+ v.ljust(w) for v, w in zip(values, widths, strict=True)
72
+ ).rstrip()
73
+
74
+ return "\n".join(
75
+ [line(header), line(["-" * w for w in widths]), *(line(row) for row in cells)]
76
+ )
77
+
78
+
79
+ def format_report(items: dict[str, Any]) -> str:
80
+ """Format the report of one process as a two-column table.
81
+
82
+ Args:
83
+ items: What :func:`info` returned.
84
+
85
+ Returns:
86
+ The table, one item per line.
87
+ """
88
+ items = dict(items)
89
+ if "rank" in items:
90
+ items["rank"] = f"{items['rank']} of {items.pop('size')}"
91
+ return _table(("item", "value"), list(items.items()))
92
+
93
+
94
+ def format_reports(reports: Sequence[dict[str, Any]]) -> str:
95
+ """Format the reports of all ranks: the common items, then one row per rank.
96
+
97
+ Args:
98
+ reports: What :func:`info` returned on each rank, in rank order.
99
+
100
+ Returns:
101
+ A table of the items that are the same on every rank, and a table with one
102
+ row per rank for the others (rank, host, local rank, launcher variable and
103
+ whatever else differs).
104
+ """
105
+ names = list(reports[0])
106
+ differs = {n for n in names if any(r.get(n) != reports[0].get(n) for r in reports)}
107
+ columns = [n for n in _RANK_COLUMNS if n in names]
108
+ columns += [n for n in names if n in differs and n not in columns]
109
+ common = [(n, reports[0][n]) for n in names if n not in columns]
110
+ rows = [[report.get(n, "-") for n in columns] for report in reports]
111
+ return f"{_table(('item', 'value'), common)}\n\n{_table(columns, rows)}"
112
+
113
+
114
+ def report(init: bool = False) -> str:
115
+ """Return the report of this process as a table (see :func:`info`)."""
116
+ return format_report(info(init))
117
+
118
+
119
+ def main(argv: list[str] | None = None) -> None:
120
+ """Print the report; the entry point of the ``maybempi`` command."""
121
+ parser = argparse.ArgumentParser(
122
+ prog="maybempi",
123
+ description="Show whether this process would use MPI, and why.",
124
+ )
125
+ parser.add_argument(
126
+ "--init",
127
+ action="store_true",
128
+ help="also start MPI (under a launcher); rank 0 then prints one table for all ranks",
129
+ )
130
+ parser.add_argument("--version", action="version", version=__version__)
131
+ args = parser.parse_args(argv)
132
+ items = info(init=args.init)
133
+ if args.init and items["size"] > 1:
134
+ reports = get_mpi().COMM_WORLD.gather(items, root=0)
135
+ if items["rank"] == 0:
136
+ print(format_reports(reports), flush=True)
137
+ return
138
+ print(format_report(items), flush=True)
139
+
140
+
141
+ if __name__ == "__main__":
142
+ main()
maybempi/launch.py ADDED
@@ -0,0 +1,204 @@
1
+ """Decide whether to use MPI, from the environment the launcher sets up.
2
+
3
+ Importing ``mpi4py.MPI`` calls ``MPI_Init``, which can take close to a second and
4
+ makes every collective cost something, even on one process. A plain
5
+ ``python script.py`` should therefore not touch MPI, even when mpi4py is
6
+ installed. :func:`launched_under_mpi` tells, without importing mpi4py, whether
7
+ the process was started by ``mpirun``/``mpiexec``/``srun``. :func:`get_mpi`
8
+ returns ``mpi4py.MPI`` then, and the serial stand-in
9
+ :class:`~maybempi.serial.SerialMPI` otherwise, so one code path serves both::
10
+
11
+ import maybempi
12
+
13
+ MPI = maybempi.get_mpi()
14
+ comm = MPI.COMM_WORLD
15
+ total = comm.allreduce(local_total, op=MPI.SUM) # local_total on one process
16
+ """
17
+
18
+ from __future__ import annotations
19
+
20
+ import os
21
+ import sys
22
+ import warnings
23
+ from typing import Any
24
+
25
+ from maybempi.serial import SerialComm, SerialMPI
26
+
27
+ __all__ = [
28
+ "LAUNCHER_VARIABLES",
29
+ "LOCAL_RANK_VARIABLES",
30
+ "OVERRIDE_VARIABLE",
31
+ "get_mpi",
32
+ "is_serial",
33
+ "launched_under_mpi",
34
+ "launcher_variable",
35
+ "local_rank",
36
+ ]
37
+
38
+ #: Per-rank variables exported by the process managers behind common launchers.
39
+ #: Each is set only for processes started *by* a launcher. ``SLURM_PROCID`` is
40
+ #: deliberately absent: it is also set for the batch script of a plain
41
+ #: ``sbatch`` job, which is not an MPI launch (``srun`` exports the PMI/PMIx
42
+ #: variables).
43
+ LAUNCHER_VARIABLES = (
44
+ "OMPI_COMM_WORLD_RANK", # Open MPI (and derivatives)
45
+ "PMI_RANK", # MPICH, Intel MPI, MS-MPI, Cray, srun --mpi=pmi2
46
+ "PMIX_RANK", # PMIx: srun --mpi=pmix, Open MPI 5
47
+ "MV2_COMM_WORLD_RANK", # MVAPICH2
48
+ "MPI_LOCALRANKID", # Hydra (mpiexec.hydra)
49
+ "ALPS_APP_PE", # Cray ALPS aprun
50
+ "PALS_RANKID", # Cray PALS
51
+ )
52
+
53
+ #: Node-local rank of the process, as exported by common launchers. They are set
54
+ #: before ``MPI_Init``, so e.g. a GPU can be chosen before MPI starts.
55
+ LOCAL_RANK_VARIABLES = (
56
+ "OMPI_COMM_WORLD_LOCAL_RANK", # Open MPI
57
+ "MV2_COMM_WORLD_LOCAL_RANK", # MVAPICH2
58
+ "MPI_LOCALRANKID", # Intel MPI, MPICH (Hydra)
59
+ "PMI_LOCAL_RANK", # MPICH / PMI
60
+ "PALS_LOCAL_RANKID", # Cray PALS
61
+ "SLURM_LOCALID", # Slurm (srun)
62
+ "LOCAL_RANK", # torchrun and others
63
+ )
64
+
65
+ #: Environment variable that forces the decision of :func:`launched_under_mpi`
66
+ #: (``1``/``true``/``yes``/``on`` or ``0``/``false``/``no``/``off``).
67
+ OVERRIDE_VARIABLE = "MAYBEMPI"
68
+
69
+ _TRUE = ("1", "true", "yes", "on")
70
+ _FALSE = ("0", "false", "no", "off")
71
+
72
+ _SERIAL_MPI = SerialMPI()
73
+ _AUTO_MPI: Any = None # the result of get_mpi(None), decided once per process
74
+
75
+
76
+ def _env_flag(name: str) -> bool | None:
77
+ """Return the boolean value of the environment variable `name`, or None if unset or unknown."""
78
+ value = os.environ.get(name)
79
+ if value is None:
80
+ return None
81
+ value = value.strip().lower()
82
+ if value in _TRUE:
83
+ return True
84
+ if value in _FALSE:
85
+ return False
86
+ return None
87
+
88
+
89
+ def local_rank() -> int:
90
+ """Return the rank of this process within its node, from the launcher's environment.
91
+
92
+ Reads the node-local rank that common launchers export (Open MPI, MVAPICH2,
93
+ Intel MPI/MPICH, PMI, Cray PALS, Slurm, ``LOCAL_RANK``, see
94
+ :data:`LOCAL_RANK_VARIABLES`). These are set before ``MPI_Init``, so this works
95
+ before MPI is initialized and without importing mpi4py.
96
+
97
+ Returns:
98
+ The node-local rank, or 0 if no launcher variable is set (a serial run).
99
+ """
100
+ for variable in LOCAL_RANK_VARIABLES:
101
+ value = os.environ.get(variable)
102
+ if value is None:
103
+ continue
104
+ try:
105
+ return int(value)
106
+ except ValueError:
107
+ continue
108
+ return 0
109
+
110
+
111
+ def launcher_variable() -> str | None:
112
+ """Return the first variable of :data:`LAUNCHER_VARIABLES` that is set, if any.
113
+
114
+ Useful to see why :func:`launched_under_mpi` decided as it did.
115
+
116
+ Returns:
117
+ The name of the variable, or None if none is set.
118
+ """
119
+ return next((v for v in LAUNCHER_VARIABLES if v in os.environ), None)
120
+
121
+
122
+ def launched_under_mpi() -> bool:
123
+ """Tell whether this process was started by an MPI launcher, without importing mpi4py.
124
+
125
+ True if a per-rank variable of a common launcher is set (Open MPI, MPICH,
126
+ Intel MPI, PMIx/``srun``, MVAPICH2, Hydra, Cray ALPS/PALS; see
127
+ :data:`LAUNCHER_VARIABLES`), or if mpi4py is already imported and MPI
128
+ initialized (using it then costs nothing more). ``MAYBEMPI=1``/``0``
129
+ overrides the detection, e.g. for a launcher whose variables are not known.
130
+
131
+ Returns:
132
+ Whether the process belongs to an MPI job.
133
+ """
134
+ override = _env_flag(OVERRIDE_VARIABLE)
135
+ if override is not None:
136
+ return override
137
+ if launcher_variable() is not None:
138
+ return True
139
+ # only look at mpi4py if the application imported it: importing it here is
140
+ # what must be avoided
141
+ mpi = sys.modules.get("mpi4py.MPI")
142
+ if mpi is not None:
143
+ try:
144
+ return bool(mpi.Is_initialized())
145
+ except AttributeError:
146
+ return False
147
+ return False
148
+
149
+
150
+ def get_mpi(use_mpi: bool | None = None) -> Any:
151
+ """Return ``mpi4py.MPI`` for an MPI run, else the serial stand-in.
152
+
153
+ Args:
154
+ use_mpi: ``None`` (the default) decides with :func:`launched_under_mpi`,
155
+ once per process. ``True`` imports mpi4py (``ImportError`` if it is
156
+ not installed). ``False`` returns the stand-in without importing
157
+ mpi4py.
158
+
159
+ Returns:
160
+ ``mpi4py.MPI``, or the one :class:`~maybempi.serial.SerialMPI` object,
161
+ which has the attributes of the module that a serial run needs.
162
+ :func:`is_serial` tells which one it is.
163
+
164
+ Warns:
165
+ RuntimeWarning: Launched under MPI, but mpi4py is not installed. Every
166
+ process then runs as if it were alone, as rank 0 of 1.
167
+ """
168
+ global _AUTO_MPI
169
+ if use_mpi is True:
170
+ from mpi4py import MPI # pyright: ignore[reportMissingImports]
171
+
172
+ return MPI
173
+ if use_mpi is False:
174
+ return _SERIAL_MPI
175
+ if _AUTO_MPI is None:
176
+ if launched_under_mpi():
177
+ try:
178
+ from mpi4py import MPI as mpi # pyright: ignore[reportMissingImports]
179
+ except ImportError:
180
+ warnings.warn(
181
+ "launched under an MPI launcher, but mpi4py is not installed: "
182
+ "every process runs serially as rank 0 of 1 (pip install mpi4py)",
183
+ RuntimeWarning,
184
+ stacklevel=2,
185
+ )
186
+ mpi = _SERIAL_MPI
187
+ _AUTO_MPI = mpi
188
+ else:
189
+ _AUTO_MPI = _SERIAL_MPI
190
+ return _AUTO_MPI
191
+
192
+
193
+ def is_serial(obj: Any) -> bool:
194
+ """Tell whether `obj` is the serial stand-in (the module or a communicator).
195
+
196
+ Args:
197
+ obj: What :func:`get_mpi` returned, or a communicator such as
198
+ ``MPI.COMM_WORLD``.
199
+
200
+ Returns:
201
+ True for :class:`~maybempi.serial.SerialMPI` and
202
+ :class:`~maybempi.serial.SerialComm` objects, False for mpi4py's.
203
+ """
204
+ return isinstance(obj, (SerialMPI, SerialComm))
maybempi/py.typed ADDED
File without changes