diff --git a/CHANGELOG.md b/CHANGELOG.md index 362b607..6c1850d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -41,6 +41,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - `KernelCatalog.from_package(..., include_dirs=None)` is now an explicit keyword; by default the source root of the top-level package (the directory containing it) is an include directory of every CUDA kernel, in addition to the kernel's own folder, so kernels can `#include "my_pkg/common.cuh"`. ### Added +- `xp.mpi.get_mpi(use_mpi=None)`, `launched_under_mpi()`, `SerialMPI` and `SerialComm`: MPI only when it is used. `launched_under_mpi()` tells from the launcher's environment (Open MPI, MPICH, Intel MPI, PMIx/`srun`, MVAPICH2, Hydra, Cray ALPS/PALS; `CUNUMPY_MPI=1/0` overrides), without importing mpi4py, whether the process was started by `mpirun`/`mpiexec`/`srun`. `get_mpi()` returns `mpi4py.MPI` then, and otherwise a serial stand-in whose `COMM_WORLD` is a communicator of size 1: object collectives return what one rank gets (`allreduce(x)` is `x`, `gather(x)` is `[x]`), buffer collectives copy (`Allreduce`, `Allgather(v)`, `Gather(v)`, `Scatter(v)`, `Sendrecv` to self; nothing with `IN_PLACE`), on NumPy and CuPy buffers, and unsupported methods raise `AttributeError`. The stand-in has the module attributes serial code needs (`IN_PLACE`, reduction ops, datatypes, `PROC_NULL`, `ROOT`, `COMM_NULL`, `Request.Waitall`, `Prequest.Startall`, `Wtime`, ...). It replaces feectools' `MockComm`/`MockMPI`, which returned `None` for every call (`MockComm().allreduce(5)` was `None`, `MockComm().rank` a function). Kept in `cunumpy/_mpi_serial.py` without dependencies on the rest of cunumpy, so that it could become a package of its own. - `cunumpy/morton.cuh` and `xp.morton_keys`, `morton_encode`, `morton_decode`, `morton_scales`, `MAX_MORTON_LEVELS`: Morton (Z-order) keys of 2D and 3D points (`uint64`, up to 32 and 21 bits per axis), the same in a kernel and on the host (bit for bit), for sorting particles along a space-filling curve and building quadtrees and octrees from sorted keys. - `xp.sort_by_key(keys, *arrays)`: one stable argsort of `keys` applied to every array; returns the sorted keys, the order and the sorted arrays. - `Kernel.from_folder(package, ...)`: the kernel of one kernel folder, with the options of `KernelCatalog.from_package` (`host_suffix`, `compile_host`, `dispatch`, `include_dirs`, CUDA options...). Every version in the folder is an implementation: `.py` (`"pyccel"`, compiled with `compile_host`, and `"python"`, uncompiled), `_numba.py`, `_numpy.py` and `_cuda.cu`. A folder's own `__init__.py` can declare `kernel = xp.Kernel.from_folder(__name__, ...)`, so code imports the kernel from where it is written; `from_package` now builds each kernel with it. `Kernel.implementations` lists them, `Kernel.selected(device=False)` tells which one a call runs. diff --git a/docs/source/api.md b/docs/source/api.md index 2bc43af..d6df722 100644 --- a/docs/source/api.md +++ b/docs/source/api.md @@ -429,6 +429,48 @@ mapping differs: device_id = xp.cuda.set_device_for_rank(mpi_rank) ``` +### `mpi.get_mpi(use_mpi=None)`, `mpi.launched_under_mpi()` + +Importing `mpi4py.MPI` starts MPI (`MPI_Init`), which takes time and makes +every collective cost something even on one process. `launched_under_mpi()` +tells, from the environment that `mpirun`/`mpiexec`/`srun` set up and without +importing mpi4py, whether the process belongs to an MPI job +(`CUNUMPY_MPI=1`/`0` overrides it). `get_mpi()` returns `mpi4py.MPI` then, and +otherwise the serial stand-in, so that the same code runs with and without +MPI: + +```python +MPI = xp.mpi.get_mpi() # decided once per process +comm = MPI.COMM_WORLD +comm.Allreduce(MPI.IN_PLACE, rho, op=MPI.SUM) # nothing to do on one process +n_total = comm.allreduce(n_local) # n_local itself +if isinstance(MPI, xp.mpi.SerialMPI): + ... # a serial run +``` + +`get_mpi(True)` imports mpi4py (`ImportError` if missing), `get_mpi(False)` +returns the stand-in. Launched under MPI without mpi4py installed, `get_mpi()` +warns (`RuntimeWarning`) and returns the stand-in. + +### `mpi.SerialMPI`, `mpi.SerialComm` + +The stand-in for `mpi4py.MPI` and its communicators of size 1. +`SerialComm` returns (object methods) or copies (buffer methods) what one rank +gets: `bcast`, `allreduce`, `reduce` and `scan` return their argument, +`gather`/`allgather` return `[x]`, `scatter([x])` returns `x`; +`Allreduce`/`Reduce`/`Allgather`/`Gather`/`Scatter`/`Alltoall`/`Scan` copy the +send buffer into the receive buffer (nothing with `IN_PLACE`), the vector +forms (`Allgatherv`, `Gatherv`, `Scatterv`) use the displacement of rank 0, +`Bcast` and `Barrier` do nothing, and `sendrecv`/`Sendrecv` work to and from +rank 0 or `PROC_NULL`. Buffers are NumPy or CuPy arrays, or mpi4py buffer +specifications (`[array, MPI.DOUBLE]`). Non-blocking versions return completed +requests. Any other method raises `AttributeError`, so a missing feature shows +instead of doing nothing. `SerialMPI` has `COMM_WORLD`, `COMM_SELF`, +`COMM_NULL` (false), `IN_PLACE`, the reduction operations, the common +datatypes (`isinstance(MPI.DOUBLE, MPI.Datatype)` holds), `PROC_NULL`, `ROOT`, +`ANY_SOURCE`, `UNDEFINED`, `Comm`/`Intracomm`, `Request`, `Prequest`, `Status`, +`Wtime()` and `Is_initialized()` (False). + ### `mpi.local_rank()` The rank of the process within its node, read from the environment variables diff --git a/docs/source/guides/mpi.md b/docs/source/guides/mpi.md index f64d54a..c355fd7 100644 --- a/docs/source/guides/mpi.md +++ b/docs/source/guides/mpi.md @@ -51,6 +51,17 @@ rank and assumes ranks are numbered contiguously per node. Prefer Importing `mpi4py.MPI` initializes MPI by default. Do it after step 1. +To run the same program serially without starting MPI, get the module from +`xp.mpi.get_mpi()` instead of importing it: it returns `mpi4py.MPI` when the +process was started by `mpirun`/`mpiexec`/`srun`, and otherwise a serial +stand-in whose `COMM_WORLD` has size 1 (see the +[API reference](../api.md)): + +```python +MPI = xp.mpi.get_mpi() +comm = MPI.COMM_WORLD +``` + ### 3. Check that MPI is CUDA-aware Only an MPI library built with CUDA support can send and receive CuPy arrays diff --git a/src/cunumpy/LLM_GUIDE.md b/src/cunumpy/LLM_GUIDE.md index d07b036..a94cbe7 100644 --- a/src/cunumpy/LLM_GUIDE.md +++ b/src/cunumpy/LLM_GUIDE.md @@ -142,6 +142,10 @@ Only transfers through cunumpy are counted (not raw `cupy.asarray`, `.get()`, MPI, accumulation and versions: ```python +MPI = ( + xp.mpi.get_mpi() +) # mpi4py.MPI under mpirun/srun, else a serial stand-in (no MPI_Init) +xp.mpi.launched_under_mpi() # from the launcher env, without importing mpi4py xp.mpi.mpi_is_cuda_aware(comm) # collective, once at startup; remembered with xp.mpi.mpi_buffer(a) as buf: comm.Send(buf, ...) # host array, CUDA-aware device diff --git a/src/cunumpy/__init__.py b/src/cunumpy/__init__.py index 17c4056..f8f94c7 100644 --- a/src/cunumpy/__init__.py +++ b/src/cunumpy/__init__.py @@ -41,7 +41,19 @@ **dict.fromkeys(kernels.__all__, "kernels"), **dict.fromkeys(rng.__all__, "rng"), **dict.fromkeys(algorithms.__all__, "algorithms"), - **dict.fromkeys(mpi.__all__, "mpi"), + # the names of cunumpy.mpi that were at the top level (not the later ones) + **dict.fromkeys( + ( + "get_mpi_cuda_aware", + "local_rank", + "mpi_buffer", + "mpi_is_cuda_aware", + "require_cuda_aware_mpi", + "set_mpi_cuda_aware", + "synchronize_for_mpi", + ), + "mpi", + ), **dict.fromkeys(profiling.__all__, "profiling"), **dict.fromkeys(memory.__all__, "memory"), "petsc_vec": "petsc", diff --git a/src/cunumpy/_mpi.py b/src/cunumpy/_mpi.py index c81288d..53f03ca 100644 --- a/src/cunumpy/_mpi.py +++ b/src/cunumpy/_mpi.py @@ -3,7 +3,6 @@ from __future__ import annotations import logging -import os from collections.abc import Generator from contextlib import contextmanager from typing import Any @@ -11,6 +10,7 @@ import array_api_compat import array_api_compat.numpy as np +from ._mpi_serial import _LOCAL_RANK_VARIABLES, local_rank # noqa: F401 - re-exported from ._transfers import _ACTIVE as _COUNTERS from ._transfers import _describe, _record from .xp import array_backend, cupy_available, to_numpy @@ -18,38 +18,6 @@ _logger = logging.getLogger(__name__) -# Node-local rank of the process, as exported by common MPI launchers. They are -# set before ``MPI_Init``, so the device can be chosen before MPI starts. -_LOCAL_RANK_VARIABLES = ( - "OMPI_COMM_WORLD_LOCAL_RANK", # Open MPI - "MV2_COMM_WORLD_LOCAL_RANK", # MVAPICH2 - "MPI_LOCALRANKID", # Intel MPI, MPICH (Hydra) - "PMI_LOCAL_RANK", # MPICH / PMI - "PALS_LOCAL_RANKID", # Cray PALS - "SLURM_LOCALID", # Slurm (srun) - "LOCAL_RANK", # torchrun and others -) - - -def local_rank() -> int: - """Rank of this process within its node, from the MPI launcher's environment. - - Reads the node-local rank that common launchers export (Open MPI, MVAPICH2, - Intel MPI/MPICH, PMI, Cray PALS, Slurm, ``LOCAL_RANK``). These variables are - set before ``MPI_Init``, so this works before MPI is initialized, and - without importing ``mpi4py``. Returns 0 if none is set (e.g. a serial run). - """ - for variable in _LOCAL_RANK_VARIABLES: - value = os.environ.get(variable) - if value is None: - continue - try: - return int(value) - except ValueError: - continue - return 0 - - def synchronize_for_mpi(*arrays: Any) -> None: """Wait for pending device work before MPI reads or writes `arrays`. diff --git a/src/cunumpy/_mpi_serial.py b/src/cunumpy/_mpi_serial.py new file mode 100644 index 0000000..88f1473 --- /dev/null +++ b/src/cunumpy/_mpi_serial.py @@ -0,0 +1,648 @@ +"""MPI only when launched under MPI, and a serial stand-in otherwise (see :mod:`cunumpy.mpi`). + +Importing ``mpi4py.MPI`` calls ``MPI_Init``, which can take close to a second +and makes every collective cost something, even on one process. A plain +``python script.py`` should therefore not touch MPI, even when mpi4py is +installed. :func:`launched_under_mpi` tells, from the environment the launcher +sets up and without importing mpi4py, whether the process was started by +``mpirun``/``mpiexec``/``srun``; :func:`get_mpi` returns ``mpi4py.MPI`` then, +and :class:`SerialMPI` otherwise, so that one code path serves both:: + + MPI = xp.mpi.get_mpi() + comm = MPI.COMM_WORLD + comm.Allreduce(MPI.IN_PLACE, rho, op=MPI.SUM) # a no-op on one process + total = comm.allreduce(local_total) # local_total itself + +:class:`SerialComm` behaves like a communicator of size 1: collectives return +or copy what they would on one rank (never ``None`` in place of a value), and +methods it does not implement raise ``AttributeError`` instead of silently +doing nothing. + +This module does not depend on the rest of cunumpy (arrays are handled by duck +typing), so that it could become a package of its own. +""" + +from __future__ import annotations + +import os +import socket +import sys +import time +import warnings +from types import MappingProxyType +from typing import Any + +# Per-rank variables exported by the process managers behind common launchers. +# Each is set only for processes started *by* a launcher. SLURM_PROCID is +# deliberately absent: it is also set for the batch script of a plain `sbatch` +# job, which is not an MPI launch (`srun` exports the PMI/PMIx variables). +_LAUNCHER_VARIABLES = ( + "OMPI_COMM_WORLD_RANK", # Open MPI (and derivatives) + "PMI_RANK", # MPICH, Intel MPI, MS-MPI, Cray, srun --mpi=pmi2 + "PMIX_RANK", # PMIx: srun --mpi=pmix, Open MPI 5 + "MV2_COMM_WORLD_RANK", # MVAPICH2 + "MPI_LOCALRANKID", # Hydra (mpiexec.hydra) + "ALPS_APP_PE", # Cray ALPS aprun + "PALS_RANKID", # Cray PALS +) + +# Node-local rank of the process, as exported by common MPI launchers. They are +# set before ``MPI_Init``, so the device can be chosen before MPI starts. +_LOCAL_RANK_VARIABLES = ( + "OMPI_COMM_WORLD_LOCAL_RANK", # Open MPI + "MV2_COMM_WORLD_LOCAL_RANK", # MVAPICH2 + "MPI_LOCALRANKID", # Intel MPI, MPICH (Hydra) + "PMI_LOCAL_RANK", # MPICH / PMI + "PALS_LOCAL_RANKID", # Cray PALS + "SLURM_LOCALID", # Slurm (srun) + "LOCAL_RANK", # torchrun and others +) + +#: Environment variable that forces the decision of :func:`launched_under_mpi` +#: (``1``/``true``/``yes``/``on`` or ``0``/``false``/``no``/``off``). +OVERRIDE_VARIABLE = "CUNUMPY_MPI" + +_TRUE = ("1", "true", "yes", "on") +_FALSE = ("0", "false", "no", "off") + + +def _env_flag(name: str) -> bool | None: + """The boolean value of the environment variable `name`, or None if unset/unknown.""" + value = os.environ.get(name) + if value is None: + return None + value = value.strip().lower() + if value in _TRUE: + return True + if value in _FALSE: + return False + return None + + +def local_rank() -> int: + """Rank of this process within its node, from the MPI launcher's environment. + + Reads the node-local rank that common launchers export (Open MPI, MVAPICH2, + Intel MPI/MPICH, PMI, Cray PALS, Slurm, ``LOCAL_RANK``). These variables are + set before ``MPI_Init``, so this works before MPI is initialized, and + without importing ``mpi4py``. Returns 0 if none is set (e.g. a serial run). + """ + for variable in _LOCAL_RANK_VARIABLES: + value = os.environ.get(variable) + if value is None: + continue + try: + return int(value) + except ValueError: + continue + return 0 + + +def launched_under_mpi() -> bool: + """Whether this process was started by an MPI launcher (without importing mpi4py). + + True if a per-rank variable of a common launcher is set (Open MPI, MPICH, + Intel MPI, PMIx/``srun``, MVAPICH2, Hydra, Cray ALPS/PALS), or if mpi4py is + already imported and MPI initialized (then using it costs nothing more). + ``CUNUMPY_MPI=1``/``0`` overrides the detection, e.g. for a launcher whose + variables are not known here. + """ + override = _env_flag(OVERRIDE_VARIABLE) + if override is not None: + return override + if any(variable in os.environ for variable in _LAUNCHER_VARIABLES): + return True + # only look at mpi4py if the application imported it: importing it here + # is what must be avoided + mpi = sys.modules.get("mpi4py.MPI") + if mpi is not None: + try: + return bool(mpi.Is_initialized()) + except AttributeError: + return False + return False + + +_AUTO_MPI: Any = None # the result of get_mpi(None), decided once + + +def get_mpi(use_mpi: bool | None = None) -> Any: + """``mpi4py.MPI`` for an MPI run, else the serial stand-in (a :class:`SerialMPI`). + + Parameters + ---------- + use_mpi : bool | None + ``None`` (the default) decides with :func:`launched_under_mpi`, once + per process; ``True`` imports mpi4py (``ImportError`` if it is not + installed); ``False`` returns the stand-in without importing it. + + Returns + ------- + module or SerialMPI + ``mpi4py.MPI``, or the one :class:`SerialMPI` object, which has the + attributes of the module that a serial run needs. + ``isinstance(MPI, xp.mpi.SerialMPI)`` tells which one it is. + + Warns + ----- + RuntimeWarning + Launched under MPI but mpi4py is not installed: every rank then runs + as if it were alone, with the same rank 0. + """ + global _AUTO_MPI + if use_mpi is True: + from mpi4py import MPI + + return MPI + if use_mpi is False: + return _SERIAL_MPI + if _AUTO_MPI is None: + if launched_under_mpi(): + try: + from mpi4py import MPI + except ImportError: + warnings.warn( + "launched under an MPI launcher, but mpi4py is not installed: " + "every process runs serially as rank 0 of 1 (pip install mpi4py)", + RuntimeWarning, + stacklevel=2, + ) + MPI = _SERIAL_MPI + _AUTO_MPI = MPI + else: + _AUTO_MPI = _SERIAL_MPI + return _AUTO_MPI + + +class _Constant: + """A named placeholder for an MPI constant (an op, a datatype, ``IN_PLACE``, ...).""" + + __slots__ = ("name",) + + def __init__(self, name: str) -> None: + self.name = name + + def __repr__(self) -> str: + return f"SerialMPI.{self.name}" + + +class _Datatype(_Constant): + """Placeholder for an MPI datatype (``isinstance(t, MPI.Datatype)`` holds).""" + + __slots__ = () + + +class _Op(_Constant): + """Placeholder for an MPI reduction operation.""" + + __slots__ = () + + +class _Null(_Constant): + """A null handle (``COMM_NULL``, ``DATATYPE_NULL``, ...): false, like mpi4py's.""" + + __slots__ = () + + def __bool__(self) -> bool: + return False + + +_IN_PLACE = _Constant("IN_PLACE") +_COMM_NULL = _Null("COMM_NULL") + + +def _buffer(spec: Any) -> Any: + """The array of an mpi4py buffer specification (``buf`` or ``[buf, ...]``).""" + if isinstance(spec, (list, tuple)): + return spec[0] + return spec + + +def _displacement(spec: Any) -> int: + """The displacement of rank 0 in a vector buffer spec ``[buf, counts, displs, type]``.""" + if isinstance(spec, (list, tuple)) and len(spec) >= 3: + displs = spec[2] + if displs is not None and not isinstance(displs, _Constant): + return int(displs[0]) + return 0 + + +def _copy(source: Any, target: Any, offset: int = 0) -> None: + """Copy the elements of the array `source` into `target`, starting at `offset`. + + Both are flattened (C order); NumPy and CuPy arrays mix (a device source is + copied to the host with ``.get()`` for a host target). + """ + if source is _IN_PLACE or source is None or target is None: + return + if hasattr(source, "get") and not hasattr(target, "get"): + source = source.get() + if not target.flags.c_contiguous: + raise ValueError("the receive buffer must be C-contiguous") + flat = target.reshape(-1) + source = source.reshape(-1) + if offset + source.size > flat.size: + raise ValueError( + f"receive buffer too small: {flat.size} elements for {source.size} " + f"at offset {offset}" + ) + flat[offset : offset + source.size] = source + + +def _check_rank(rank: int, what: str) -> None: + if rank not in (0, SerialMPI.PROC_NULL, SerialMPI.ANY_SOURCE): + raise ValueError(f"{what}={rank}: a serial communicator has only rank 0") + + +class SerialRequest: + """A completed request, returned by the non-blocking calls of :class:`SerialComm`.""" + + def __init__(self, result: Any = None) -> None: + self._result = result + + def Wait(self, status: Any = None) -> None: + return None + + def Test(self, status: Any = None) -> bool: + return True + + def wait(self, status: Any = None) -> Any: + return self._result + + def test(self, status: Any = None) -> tuple[bool, Any]: + return True, self._result + + def Free(self) -> None: + return None + + def Cancel(self) -> None: + return None + + @staticmethod + def Waitall(requests: Any, statuses: Any = None) -> None: + return None + + @staticmethod + def waitall(requests: Any, statuses: Any = None) -> list[Any]: + return [request.wait() for request in requests] + + @staticmethod + def Testall(requests: Any, statuses: Any = None) -> bool: + return True + + @staticmethod + def Waitany(requests: Any, status: Any = None) -> int: + return 0 if requests else SerialMPI.UNDEFINED + + +class SerialPrequest(SerialRequest): + """A persistent request (``Send_init``/``Recv_init``), for ``Startall``/``Waitall``.""" + + def Start(self) -> None: + return None + + @staticmethod + def Startall(requests: Any) -> None: + return None + + +class SerialComm: + """A communicator of size 1, with the mpi4py ``Comm`` methods a serial run needs. + + Collectives return (object methods) or copy (buffer methods) what they + would on one rank: ``allreduce(x)`` is ``x``, ``gather(x)`` is ``[x]``, + ``Allreduce(send, recv)`` copies `send` into `recv` (nothing with + ``IN_PLACE``), ``Bcast`` does nothing. Point-to-point calls are only + supported to and from rank 0 itself (``sendrecv``, ``Sendrecv``) or + ``PROC_NULL``. Buffers may be NumPy or CuPy arrays, or mpi4py buffer specs + (``[array, MPI.DOUBLE]``). Other methods raise ``AttributeError``. + """ + + rank = 0 + size = 1 + + def __init__(self, name: str = "COMM_WORLD") -> None: + self._name = name + + def __repr__(self) -> str: + return f"SerialComm({self._name})" + + # ---------------------------------------------------------------- queries + def Get_rank(self) -> int: + return 0 + + def Get_size(self) -> int: + return 1 + + def Get_name(self) -> str: + return self._name + + def Is_inter(self) -> bool: + return False + + def Is_intra(self) -> bool: + return True + + # -------------------------------------------------- communicator creation + def Dup(self, info: Any = None) -> SerialComm: + return SerialComm(self._name) + + Clone = Dup + + def Split(self, color: int = 0, key: int = 0) -> Any: + if color == SerialMPI.UNDEFINED: + return SerialMPI.COMM_NULL + return SerialComm(self._name) + + def Free(self) -> None: + return None + + def Abort(self, errorcode: int = 0) -> None: + raise SystemExit(errorcode) + + # ------------------------------------------------------ synchronization + def Barrier(self) -> None: + return None + + barrier = Barrier + + def Ibarrier(self) -> SerialRequest: + return SerialRequest() + + # ------------------------------------------------- collectives, objects + def bcast(self, obj: Any, root: int = 0) -> Any: + _check_rank(root, "root") + return obj + + def reduce(self, sendobj: Any, op: Any = None, root: int = 0) -> Any: + _check_rank(root, "root") + return sendobj + + def allreduce(self, sendobj: Any, op: Any = None) -> Any: + return sendobj + + def scan(self, sendobj: Any, op: Any = None) -> Any: + return sendobj + + def exscan(self, sendobj: Any, op: Any = None) -> None: + return None # undefined on rank 0, None in mpi4py + + def gather(self, sendobj: Any, root: int = 0) -> list[Any]: + _check_rank(root, "root") + return [sendobj] + + def allgather(self, sendobj: Any) -> list[Any]: + return [sendobj] + + def scatter(self, sendobj: Any, root: int = 0) -> Any: + _check_rank(root, "root") + items = list(sendobj) + if len(items) != 1: + raise ValueError(f"scatter on 1 process needs 1 item, got {len(items)}") + return items[0] + + def alltoall(self, sendobj: Any) -> list[Any]: + items = list(sendobj) + if len(items) != 1: + raise ValueError(f"alltoall on 1 process needs 1 item, got {len(items)}") + return items + + def sendrecv( + self, + sendobj: Any, + dest: int = 0, + sendtag: int = 0, + recvbuf: Any = None, + source: int = 0, + recvtag: int = 0, + status: Any = None, + ) -> Any: + _check_rank(dest, "dest") + _check_rank(source, "source") + if source == SerialMPI.PROC_NULL: + return None + return sendobj if dest != SerialMPI.PROC_NULL else None + + def ibcast(self, obj: Any, root: int = 0) -> SerialRequest: + return SerialRequest(self.bcast(obj, root)) + + def iallreduce(self, sendobj: Any, op: Any = None) -> SerialRequest: + return SerialRequest(sendobj) + + # ------------------------------------------------- collectives, buffers + def Bcast(self, buf: Any, root: int = 0) -> None: + _check_rank(root, "root") + + def Reduce(self, sendbuf: Any, recvbuf: Any, op: Any = None, root: int = 0) -> None: + _check_rank(root, "root") + _copy(_buffer(sendbuf), _buffer(recvbuf)) + + def Allreduce(self, sendbuf: Any, recvbuf: Any, op: Any = None) -> None: + _copy(_buffer(sendbuf), _buffer(recvbuf)) + + def Scan(self, sendbuf: Any, recvbuf: Any, op: Any = None) -> None: + _copy(_buffer(sendbuf), _buffer(recvbuf)) + + def Exscan(self, sendbuf: Any, recvbuf: Any, op: Any = None) -> None: + return None # the receive buffer of rank 0 is undefined + + def Gather(self, sendbuf: Any, recvbuf: Any, root: int = 0) -> None: + _check_rank(root, "root") + _copy(_buffer(sendbuf), _buffer(recvbuf)) + + def Gatherv(self, sendbuf: Any, recvbuf: Any, root: int = 0) -> None: + _check_rank(root, "root") + _copy(_buffer(sendbuf), _buffer(recvbuf), _displacement(recvbuf)) + + def Allgather(self, sendbuf: Any, recvbuf: Any) -> None: + _copy(_buffer(sendbuf), _buffer(recvbuf)) + + def Allgatherv(self, sendbuf: Any, recvbuf: Any) -> None: + _copy(_buffer(sendbuf), _buffer(recvbuf), _displacement(recvbuf)) + + def Scatter(self, sendbuf: Any, recvbuf: Any, root: int = 0) -> None: + _check_rank(root, "root") + if recvbuf is not _IN_PLACE: + _copy(_buffer(sendbuf), _buffer(recvbuf)) + + def Scatterv(self, sendbuf: Any, recvbuf: Any, root: int = 0) -> None: + _check_rank(root, "root") + if recvbuf is _IN_PLACE: + return + source = _buffer(sendbuf).reshape(-1) + start = _displacement(sendbuf) + target = _buffer(recvbuf) + _copy(source[start : start + target.size], target) + + def Alltoall(self, sendbuf: Any, recvbuf: Any) -> None: + _copy(_buffer(sendbuf), _buffer(recvbuf)) + + def Sendrecv( + self, + sendbuf: Any, + dest: int = 0, + sendtag: int = 0, + recvbuf: Any = None, + source: int = 0, + recvtag: int = 0, + status: Any = None, + ) -> None: + _check_rank(dest, "dest") + _check_rank(source, "source") + if SerialMPI.PROC_NULL in (dest, source): + return + _copy(_buffer(sendbuf), _buffer(recvbuf)) + + def Ibcast(self, buf: Any, root: int = 0) -> SerialRequest: + self.Bcast(buf, root) + return SerialRequest() + + def Iallreduce(self, sendbuf: Any, recvbuf: Any, op: Any = None) -> SerialRequest: + self.Allreduce(sendbuf, recvbuf, op) + return SerialRequest() + + def Iallgather(self, sendbuf: Any, recvbuf: Any) -> SerialRequest: + self.Allgather(sendbuf, recvbuf) + return SerialRequest() + + +_COMM_WORLD = SerialComm("COMM_WORLD") +_COMM_SELF = SerialComm("COMM_SELF") + + +class SerialStatus: + """Stand-in for ``MPI.Status`` (source 0, tag 0).""" + + source = 0 + tag = 0 + error = 0 + + def Get_source(self) -> int: + return 0 + + def Get_tag(self) -> int: + return 0 + + def Get_count(self, datatype: Any = None) -> int: + return 0 + + +class SerialMPI: + """Stand-in for the ``mpi4py.MPI`` module in a serial run (see :func:`get_mpi`). + + :func:`get_mpi` returns one instance; ``isinstance(MPI, SerialMPI)`` tells a + serial run from an MPI one. ``COMM_WORLD`` and ``COMM_SELF`` are :class:`SerialComm` objects; the + reduction operations, datatypes and other constants are placeholders that + :class:`SerialComm` accepts. ``Is_initialized()`` is False: MPI itself is + never started. + """ + + COMM_WORLD = _COMM_WORLD + COMM_SELF = _COMM_SELF + COMM_NULL = _COMM_NULL + Comm = Intracomm = SerialComm + Request = SerialRequest + Status = SerialStatus + + IN_PLACE = _IN_PLACE + BOTTOM = _Constant("BOTTOM") + DATATYPE_NULL = _Null("DATATYPE_NULL") + REQUEST_NULL = _Null("REQUEST_NULL") + OP_NULL = _Null("OP_NULL") + PROC_NULL = -2 + ANY_SOURCE = -1 + ANY_TAG = -1 + ROOT = -3 + UNDEFINED = -32766 + SUCCESS = 0 + + Datatype = _Datatype + Op = _Op + Prequest = SerialPrequest + + # reduction operations + SUM = _Op("SUM") + PROD = _Op("PROD") + MAX = _Op("MAX") + MIN = _Op("MIN") + LAND = _Op("LAND") + LOR = _Op("LOR") + LXOR = _Op("LXOR") + BAND = _Op("BAND") + BOR = _Op("BOR") + BXOR = _Op("BXOR") + MAXLOC = _Op("MAXLOC") + MINLOC = _Op("MINLOC") + REPLACE = _Op("REPLACE") + + # datatypes + BYTE = _Datatype("BYTE") + CHAR = _Datatype("CHAR") + BOOL = _Datatype("BOOL") + C_BOOL = _Datatype("C_BOOL") + INT = _Datatype("INT") + LONG = _Datatype("LONG") + LONG_LONG = _Datatype("LONG_LONG") + UNSIGNED = _Datatype("UNSIGNED") + UNSIGNED_LONG = _Datatype("UNSIGNED_LONG") + INT8_T = _Datatype("INT8_T") + INT16_T = _Datatype("INT16_T") + INT32_T = _Datatype("INT32_T") + INT64_T = _Datatype("INT64_T") + UINT8_T = _Datatype("UINT8_T") + UINT16_T = _Datatype("UINT16_T") + UINT32_T = _Datatype("UINT32_T") + UINT64_T = _Datatype("UINT64_T") + FLOAT = _Datatype("FLOAT") + DOUBLE = _Datatype("DOUBLE") + LONG_DOUBLE = _Datatype("LONG_DOUBLE") + C_FLOAT_COMPLEX = _Datatype("C_FLOAT_COMPLEX") + C_DOUBLE_COMPLEX = _Datatype("C_DOUBLE_COMPLEX") + COMPLEX = _Datatype("COMPLEX") + DOUBLE_COMPLEX = _Datatype("DOUBLE_COMPLEX") + + # NumPy type characters to datatypes, like mpi4py's (private) MPI._typedict + _typedict = MappingProxyType({ + "b": INT8_T, "h": INT16_T, "i": INT32_T, "l": LONG, "q": INT64_T, + "B": UINT8_T, "H": UINT16_T, "I": UINT32_T, "L": UNSIGNED_LONG, "Q": UINT64_T, + "f": FLOAT, "d": DOUBLE, "g": LONG_DOUBLE, "?": C_BOOL, + "F": C_FLOAT_COMPLEX, "D": C_DOUBLE_COMPLEX, + }) # fmt: skip + + def __repr__(self) -> str: + return "" + + @staticmethod + def Wtime() -> float: + return time.time() + + @staticmethod + def Wtick() -> float: + return time.get_clock_info("time").resolution + + @staticmethod + def Is_initialized() -> bool: + return False + + @staticmethod + def Is_finalized() -> bool: + return False + + @staticmethod + def Init() -> None: + return None + + @staticmethod + def Finalize() -> None: + return None + + @staticmethod + def Get_processor_name() -> str: + return socket.gethostname() + + @staticmethod + def Query_thread() -> int: + return 0 + + +_SERIAL_MPI = SerialMPI() diff --git a/src/cunumpy/mpi.py b/src/cunumpy/mpi.py index 83955a4..da454cc 100644 --- a/src/cunumpy/mpi.py +++ b/src/cunumpy/mpi.py @@ -1,30 +1,54 @@ -"""MPI with NumPy or CuPy arrays. +"""MPI with NumPy or CuPy arrays, and serial runs without MPI. -:func:`mpi_buffer` hands an array to mpi4py: the device array itself when the -MPI library is CUDA-aware, a host copy otherwise; :func:`mpi_is_cuda_aware` -finds out which. :func:`local_rank` is the rank of this process on its node -(to pick a GPU, see :func:`cunumpy.cuda.bind_local_device`):: +:func:`get_mpi` returns ``mpi4py.MPI`` when the process was started by an MPI +launcher (:func:`launched_under_mpi`, decided from the environment without +importing mpi4py), and :class:`SerialMPI` otherwise: a stand-in whose +``COMM_WORLD`` is a :class:`SerialComm` of size 1, so that the same code runs +serially without starting MPI:: import cunumpy as xp + MPI = xp.mpi.get_mpi() + comm = MPI.COMM_WORLD with xp.mpi.mpi_buffer(rho, send=True, recv=True) as buf: comm.Allreduce(MPI.IN_PLACE, buf, op=MPI.SUM) +:func:`mpi_buffer` hands an array to mpi4py: the device array itself when the +MPI library is CUDA-aware, a host copy otherwise; :func:`mpi_is_cuda_aware` +finds out which. :func:`local_rank` is the rank of this process on its node +(to pick a GPU, see :func:`cunumpy.cuda.bind_local_device`). + mpi4py is imported only by the functions that need it. """ from ._mpi import ( get_mpi_cuda_aware, - local_rank, mpi_buffer, mpi_is_cuda_aware, require_cuda_aware_mpi, set_mpi_cuda_aware, synchronize_for_mpi, ) +from ._mpi_serial import ( + OVERRIDE_VARIABLE, + SerialComm, + SerialMPI, + SerialRequest, + SerialStatus, + get_mpi, + launched_under_mpi, + local_rank, +) __all__ = [ + "OVERRIDE_VARIABLE", + "SerialComm", + "SerialMPI", + "SerialRequest", + "SerialStatus", + "get_mpi", "get_mpi_cuda_aware", + "launched_under_mpi", "local_rank", "mpi_buffer", "mpi_is_cuda_aware", diff --git a/tests/unit/test_mpi_serial.py b/tests/unit/test_mpi_serial.py new file mode 100644 index 0000000..e233c1a --- /dev/null +++ b/tests/unit/test_mpi_serial.py @@ -0,0 +1,301 @@ +"""Tests for `xp.mpi.get_mpi`, `launched_under_mpi` and the serial stand-in `SerialComm`.""" + +import sys +import types + +import numpy as np +import pytest + +import cunumpy as xp +from cunumpy import _mpi_serial +from cunumpy.mpi import SerialMPI, get_mpi, launched_under_mpi + +MPI = get_mpi(False) +comm = MPI.COMM_WORLD + + +@pytest.fixture +def clean_env(monkeypatch): + """No launcher variables, no override, no mpi4py imported, no cached decision.""" + for variable in (*_mpi_serial._LAUNCHER_VARIABLES, _mpi_serial.OVERRIDE_VARIABLE): + monkeypatch.delenv(variable, raising=False) + monkeypatch.delitem(sys.modules, "mpi4py.MPI", raising=False) + monkeypatch.setattr(_mpi_serial, "_AUTO_MPI", None) + return monkeypatch + + +@pytest.fixture +def fake_mpi4py(clean_env): + """A stand-in for an installed mpi4py (without starting MPI).""" + module = types.ModuleType("mpi4py.MPI") + module.Is_initialized = lambda: True + package = types.ModuleType("mpi4py") + package.MPI = module + clean_env.setitem(sys.modules, "mpi4py", package) + clean_env.setitem(sys.modules, "mpi4py.MPI", module) + return module + + +@pytest.fixture +def no_mpi4py(clean_env): + clean_env.setitem(sys.modules, "mpi4py", None) + clean_env.setitem(sys.modules, "mpi4py.MPI", None) + return clean_env + + +# ------------------------------------------------------------ the decision + + +def test_serial_by_default(clean_env): + assert launched_under_mpi() is False + assert isinstance(get_mpi(), SerialMPI) + + +@pytest.mark.parametrize("variable", _mpi_serial._LAUNCHER_VARIABLES) +def test_launcher_variables(clean_env, variable): + clean_env.setenv(variable, "3") + assert launched_under_mpi() is True + + +def test_slurm_batch_script_is_not_an_mpi_launch(clean_env): + clean_env.setenv("SLURM_PROCID", "0") + assert launched_under_mpi() is False + + +@pytest.mark.parametrize( + ("value", "launcher", "expected"), + [("1", False, True), ("on", False, True), ("0", True, False), ("no", True, False), + ("maybe", True, True), ("maybe", False, False)], +) # fmt: skip +def test_override(clean_env, value, launcher, expected): + clean_env.setenv("CUNUMPY_MPI", value) + if launcher: + clean_env.setenv("PMI_RANK", "0") + assert launched_under_mpi() is expected + + +def test_initialized_mpi4py_counts_as_mpi(fake_mpi4py): + assert launched_under_mpi() is True + fake_mpi4py.Is_initialized = lambda: False + assert launched_under_mpi() is False + + +def test_get_mpi_under_a_launcher(fake_mpi4py, clean_env): + fake_mpi4py.Is_initialized = lambda: False + clean_env.setenv("OMPI_COMM_WORLD_RANK", "0") + assert get_mpi() is fake_mpi4py + clean_env.delenv("OMPI_COMM_WORLD_RANK") + assert get_mpi() is fake_mpi4py # decided once per process + + +def test_get_mpi_explicit(fake_mpi4py): + assert get_mpi(True) is fake_mpi4py + assert get_mpi(False) is MPI + + +def test_launcher_without_mpi4py_warns(no_mpi4py): + no_mpi4py.setenv("PMI_RANK", "0") + with pytest.warns(RuntimeWarning, match="mpi4py is not installed"): + assert isinstance(get_mpi(), SerialMPI) + + +def test_explicit_mpi_without_mpi4py_raises(no_mpi4py): + with pytest.raises(ImportError): + get_mpi(True) + + +def test_exported(): + assert xp.mpi.get_mpi is get_mpi + assert xp.mpi.local_rank is _mpi_serial.local_rank + assert "get_mpi" not in xp._MOVED # never was at the top level + + +# -------------------------------------------------- the serial communicator + + +def test_queries(): + assert comm.rank == 0 and comm.size == 1 + assert comm.Get_rank() == 0 and comm.Get_size() == 1 + assert isinstance(comm, MPI.Comm) and isinstance(comm, MPI.Intracomm) + assert comm.Is_intra() and not comm.Is_inter() + assert MPI.COMM_SELF.Get_size() == 1 + assert repr(comm) == "SerialComm(COMM_WORLD)" + + +def test_object_collectives_return_the_value(): + value = {"a": 1} + assert comm.bcast(value) is value + assert comm.allreduce(5, op=MPI.SUM) == 5 + assert comm.allreduce(2.5, op=MPI.MAX) == 2.5 + assert comm.reduce(7) == 7 + assert comm.scan(3) == 3 + assert comm.exscan(3) is None # as mpi4py on rank 0 + assert comm.gather(value) == [value] + assert comm.allgather(value) == [value] + assert comm.scatter([value]) is value + assert comm.alltoall([value]) == [value] + assert comm.sendrecv(value, dest=0, source=0) is value + assert comm.ibcast(value).wait() is value + assert comm.iallreduce(4).wait() == 4 + assert comm.Barrier() is None and comm.barrier() is None + + +def test_object_collectives_check_sizes_and_ranks(): + with pytest.raises(ValueError, match="1 item"): + comm.scatter([1, 2]) + with pytest.raises(ValueError, match="only rank 0"): + comm.bcast(1, root=1) + with pytest.raises(ValueError, match="only rank 0"): + comm.sendrecv(1, dest=1) + assert comm.sendrecv(1, dest=MPI.PROC_NULL, source=MPI.PROC_NULL) is None + + +def test_unknown_methods_raise(): + # MockComm returned None for everything, which hid missing support + with pytest.raises(AttributeError): + _ = comm.Create_cart + with pytest.raises(AttributeError): + _ = MPI.Win + + +def test_buffer_collectives_copy(): + send = np.arange(6.0).reshape(2, 3) + for call in ( + lambda r: comm.Allreduce(send, r, op=MPI.SUM), + lambda r: comm.Reduce(send, r, op=MPI.SUM, root=0), + lambda r: comm.Allgather(send, r), + lambda r: comm.Gather(send, r), + lambda r: comm.Scatter(send, r), + lambda r: comm.Alltoall(send, r), + lambda r: comm.Scan(send, r), + lambda r: comm.Sendrecv(send, dest=0, recvbuf=r, source=0), + lambda r: comm.Iallreduce(send, r).Wait(), + lambda r: comm.Iallgather(send, r).Wait(), + ): + recv = np.zeros_like(send) + call(recv) + np.testing.assert_array_equal(recv, send) + + +def test_buffer_specs_and_in_place(): + data = np.arange(4.0) + recv = np.zeros(4) + comm.Allreduce([data, MPI.DOUBLE], [recv, MPI.DOUBLE], op=MPI.SUM) + np.testing.assert_array_equal(recv, data) + before = data.copy() + comm.Allreduce(MPI.IN_PLACE, data, op=MPI.SUM) + comm.Bcast(data, root=0) + comm.Ibcast(data).Wait() + np.testing.assert_array_equal(data, before) + comm.Exscan(data, recv) # undefined on rank 0: untouched + np.testing.assert_array_equal(recv, before) + + +def test_vector_collectives_use_the_displacement(): + send = np.array([1.0, 2.0]) + recv = np.zeros(5) + comm.Allgatherv(send, [recv, [2], [3], MPI.DOUBLE]) + np.testing.assert_array_equal(recv, [0, 0, 0, 1, 2]) + recv[:] = 0 + comm.Gatherv(send, [recv, [2], [1]]) + np.testing.assert_array_equal(recv, [0, 1, 2, 0, 0]) + part = np.zeros(2) + comm.Scatterv([np.arange(5.0), [2], [2], MPI.DOUBLE], part) + np.testing.assert_array_equal(part, [2, 3]) + + +def test_buffer_errors(): + with pytest.raises(ValueError, match="too small"): + comm.Allreduce(np.ones(3), np.zeros(2)) + with pytest.raises(ValueError, match="C-contiguous"): + comm.Allreduce(np.ones(2), np.zeros((2, 2))[:, 0]) + with pytest.raises(ValueError, match="only rank 0"): + comm.Sendrecv(np.ones(2), dest=1, recvbuf=np.zeros(2)) + comm.Sendrecv(np.ones(2), dest=MPI.PROC_NULL, recvbuf=np.zeros(2)) + + +class _DeviceArray: + """Duck-typed device array: `.get()` copies it to the host.""" + + def __init__(self, data): + self._data = np.asarray(data) + self.flags = self._data.flags + + def get(self): + return self._data.copy() + + def reshape(self, *shape): + return _DeviceArray(self._data.reshape(*shape)) + + @property + def size(self): + return self._data.size + + +def test_device_source_to_host_target(): + recv = np.zeros(3) + comm.Allreduce(_DeviceArray([1.0, 2.0, 3.0]), recv) + np.testing.assert_array_equal(recv, [1, 2, 3]) + + +def test_split_dup_and_requests(): + assert comm.Split(0, 0).Get_size() == 1 + assert comm.Split(MPI.UNDEFINED) is MPI.COMM_NULL + assert comm.Dup().Get_rank() == 0 and comm.Clone().Get_size() == 1 + requests = [comm.Ibarrier(), comm.Ibcast(np.zeros(1))] + assert MPI.Request.Waitall(requests) is None + assert MPI.Request.Testall(requests) is True + assert MPI.Request.waitall([comm.ibcast(1)]) == [1] + status = MPI.Status() + assert status.Get_source() == 0 and status.Get_tag() == 0 + + +def test_module_functions(): + assert MPI.Is_initialized() is False and MPI.Is_finalized() is False + assert MPI.Wtime() > 0 and MPI.Wtick() > 0 + assert isinstance(MPI.Get_processor_name(), str) + assert repr(MPI.SUM) == "SerialMPI.SUM" + assert MPI.PROC_NULL == -2 and MPI.ANY_SOURCE == -1 and MPI.ROOT == -3 + assert isinstance(MPI, SerialMPI) and get_mpi(False) is MPI + + +def test_constants_used_by_struphy_and_feectools(): + assert isinstance(MPI.DOUBLE, MPI.Datatype) and not isinstance( + MPI.SUM, MPI.Datatype + ) + assert isinstance(MPI.LOR, MPI.Op) + assert MPI._typedict[np.dtype(np.float64).char] is MPI.DOUBLE + # null handles are false, communicators true, like in mpi4py + assert not MPI.COMM_NULL and not MPI.DATATYPE_NULL + assert comm and comm != MPI.COMM_NULL + MPI.Prequest.Startall([]) + assert MPI.Prequest.Waitall([]) is None + # usable in annotations evaluated at definition time + assert (MPI.Intracomm | None) is not None + + +@pytest.mark.parametrize("backend", ["numpy", pytest.param("cupy", marks=pytest.mark.skipif( + not xp.cupy_available(), reason="CuPy/GPU not available"))]) # fmt: skip +def test_buffers_on_either_backend(backend): + with xp.use_backend(backend): + send = xp.arange(4.0) + recv = xp.zeros(4) + comm.Allreduce(send, recv, op=MPI.SUM) + np.testing.assert_array_equal(xp.to_numpy(recv), np.arange(4.0)) + host = np.zeros(4) + comm.Allgather(send, host) # a device array into a host buffer + np.testing.assert_array_equal(host, np.arange(4.0)) + + +def test_matches_mpi4py_on_one_process(): + """The same calls on mpi4py's COMM_SELF give the same results (if mpi4py is there).""" + real = pytest.importorskip("mpi4py.MPI") + self_comm = real.COMM_SELF + assert self_comm.allreduce(5) == comm.allreduce(5) + assert self_comm.gather(3) == comm.gather(3) + assert self_comm.scatter([4]) == comm.scatter([4]) + assert self_comm.exscan(3) == comm.exscan(3) + send, recv_real, recv_serial = np.arange(3.0), np.zeros(3), np.zeros(3) + self_comm.Allreduce(send, recv_real, op=real.SUM) + comm.Allreduce(send, recv_serial, op=MPI.SUM) + np.testing.assert_array_equal(recv_real, recv_serial)