diff --git a/examples/distributed.py b/examples/distributed.py index 733b41f74..c8601aaba 100644 --- a/examples/distributed.py +++ b/examples/distributed.py @@ -30,6 +30,9 @@ def main(): # Find the partition outputs = pt.DictOfNamedArrays({"out": y}) + + pt.verify_distributed_dag_pre_partition(comm, outputs) + distributed_parts = find_distributed_partition(outputs) distributed_parts, _ = number_distributed_tags( comm, distributed_parts, base_tag=42) diff --git a/pytato/__init__.py b/pytato/__init__.py index 7f7a61066..d18b1a761 100644 --- a/pytato/__init__.py +++ b/pytato/__init__.py @@ -100,7 +100,8 @@ def set_debug_enabled(flag: bool) -> None: from pytato.distributed.partition import ( find_distributed_partition, DistributedGraphPart, DistributedGraphPartition) from pytato.distributed.tags import number_distributed_tags -from pytato.distributed.verify import verify_distributed_partition +from pytato.distributed.verify import (verify_distributed_partition, + verify_distributed_dag_pre_partition) from pytato.distributed.execute import execute_distributed_partition from pytato.transform.lower_to_index_lambda import to_index_lambda @@ -161,6 +162,7 @@ def set_debug_enabled(flag: bool) -> None: "number_distributed_tags", "execute_distributed_partition", "verify_distributed_partition", + "verify_distributed_dag_pre_partition", "generate_code_for_partition", diff --git a/pytato/distributed/partition.py b/pytato/distributed/partition.py index 1549d9755..f9eacc7ee 100644 --- a/pytato/distributed/partition.py +++ b/pytato/distributed/partition.py @@ -33,8 +33,10 @@ THE SOFTWARE. """ +from functools import reduce from typing import ( - Tuple, Any, Mapping, FrozenSet, Set, Dict, cast, Iterable, Callable, List) + Tuple, Any, Mapping, FrozenSet, Set, Dict, cast, Iterable, + Callable, List, TypeVar) from functools import cached_property import attrs @@ -53,10 +55,32 @@ CombineMapper) from pytato.partition import GraphPart, GraphPartition, PartId, GraphPartitioner from pytato.distributed.nodes import ( - DistributedRecv, DistributedSend, DistributedSendRefHolder) + DistributedRecv, DistributedSend, DistributedSendRefHolder, CommTagType) from pytato.analysis import DirectPredecessorsGetter +@attrs.define(frozen=True) +class CommunicationOpIdentifier: + """Identifies a communication operation (consisting of a pair of + a send and a receive). + .. attribute:: src_rank + .. attribute:: dest_rank + .. attribute:: comm_tag + .. note:: + In :func:`find_distributed_partition`, we use instances of this type as + though they identify sends or receives, i.e. just a single end of the + communication. Realize that this is only true given the additional + context of which rank is the local rank. + """ + src_rank: int + dest_rank: int + comm_tag: CommTagType + + +_KeyT = TypeVar("_KeyT") +_ValueT = TypeVar("_ValueT") + + # {{{ distributed graph partition @attrs.define(frozen=True, slots=False) @@ -85,6 +109,82 @@ class DistributedGraphPartition(GraphPartition): # }}} +# {{{ _LocalSendRecvDepGatherer + +def _send_to_comm_id( + local_rank: int, send: DistributedSend) -> CommunicationOpIdentifier: + return CommunicationOpIdentifier( + src_rank=local_rank, + dest_rank=send.dest_rank, + comm_tag=send.comm_tag) + + +def _recv_to_comm_id( + local_rank: int, recv: DistributedRecv) -> CommunicationOpIdentifier: + return CommunicationOpIdentifier( + src_rank=recv.src_rank, + dest_rank=local_rank, + comm_tag=recv.comm_tag) + + +class _LocalSendRecvDepGatherer( + CombineMapper[FrozenSet[CommunicationOpIdentifier]]): + def __init__(self, local_rank: int) -> None: + super().__init__() + self.local_send_id_to_needed_local_recv_ids: \ + Dict[CommunicationOpIdentifier, + FrozenSet[CommunicationOpIdentifier]] = {} + + self.local_recv_id_to_recv_node: \ + Dict[CommunicationOpIdentifier, DistributedRecv] = {} + self.local_send_id_to_send_node: \ + Dict[CommunicationOpIdentifier, DistributedSend] = {} + + self.local_rank = local_rank + + def combine( + self, *args: FrozenSet[CommunicationOpIdentifier] + ) -> FrozenSet[CommunicationOpIdentifier]: + return reduce(frozenset.union, args, frozenset()) + + def map_distributed_send_ref_holder(self, + expr: DistributedSendRefHolder + ) -> FrozenSet[CommunicationOpIdentifier]: + send_id = _send_to_comm_id(self.local_rank, expr.send) + + if send_id in self.local_send_id_to_needed_local_recv_ids: + raise ValueError(f"Multiple sends found for '{send_id}'") + + self.local_send_id_to_needed_local_recv_ids[send_id] = \ + self.rec(expr.send.data) + + assert send_id not in self.local_send_id_to_send_node + self.local_send_id_to_send_node[send_id] = expr.send + + return self.rec(expr.passthrough_data) + + def _map_input_base(self, expr: Array) -> FrozenSet[CommunicationOpIdentifier]: + return frozenset() + + map_placeholder = _map_input_base + map_data_wrapper = _map_input_base + map_size_param = _map_input_base + + def map_distributed_recv( + self, expr: DistributedRecv + ) -> FrozenSet[CommunicationOpIdentifier]: + recv_id = _recv_to_comm_id(self.local_rank, expr) + + if recv_id in self.local_recv_id_to_recv_node: + raise ValueError(f"Multiple receives found for '{recv_id}'") + + self.local_recv_id_to_recv_node[recv_id] = expr + + return frozenset({recv_id}) + +# }}} + + # {{{ _partition_to_distributed_partition def _map_distributed_graph_partition_nodes( diff --git a/pytato/distributed/verify.py b/pytato/distributed/verify.py index e37d8a83a..13ed1171d 100644 --- a/pytato/distributed/verify.py +++ b/pytato/distributed/verify.py @@ -30,14 +30,16 @@ """ -from typing import Any, FrozenSet, Dict, Set, Optional, Sequence, TYPE_CHECKING +from typing import (Any, FrozenSet, Dict, Set, Optional, Sequence, + TYPE_CHECKING, Mapping) import numpy as np from pytato.distributed.nodes import CommTagType, DistributedRecv from pytato.partition import PartId -from pytato.distributed.partition import DistributedGraphPartition -from pytato.array import ShapeType +from pytato.distributed.partition import (DistributedGraphPartition, + _KeyT, _ValueT, CommunicationOpIdentifier) +from pytato.array import ShapeType, DictOfNamedArrays import attrs @@ -48,6 +50,7 @@ if TYPE_CHECKING: import mpi4py.MPI + from mpi4py import MPI # {{{ data structures @@ -122,6 +125,78 @@ class MissingRecvError(DistributedPartitionVerificationError): # }}} +# {{{ _dict_union_mpi + +def _dict_union_mpi( + dict_a: Mapping[_KeyT, _ValueT], dict_b: Mapping[_KeyT, _ValueT], + mpi_data_type: MPI.Datatype) -> Mapping[_KeyT, _ValueT]: + assert mpi_data_type is None + result = dict(dict_a) + result.update(dict_b) + return result + +# }}} + + +# {{{ _get_comm_to_needed_comms + +def _get_comm_to_needed_comms(mpi_communicator: mpi4py.MPI.Comm, + outputs: DictOfNamedArrays) -> \ + Dict[CommunicationOpIdentifier, FrozenSet[CommunicationOpIdentifier]]: + my_rank = mpi_communicator.rank + + from pytato.distributed.partition import _LocalSendRecvDepGatherer + lsrdg = _LocalSendRecvDepGatherer(local_rank=my_rank) + lsrdg(outputs) + local_send_id_to_needed_local_recv_ids = \ + lsrdg.local_send_id_to_needed_local_recv_ids + + from mpi4py import MPI + dict_union_mpi_op = MPI.Op.Create( + # type ignore reason: mpi4py misdeclares op functions as returning + # None. + _dict_union_mpi, # type: ignore[arg-type] + commute=True) + try: + # FIXME: allreduce might not be necessary for all use cases + comm_to_needed_comms: \ + Dict[CommunicationOpIdentifier, FrozenSet[CommunicationOpIdentifier]] = \ + mpi_communicator.allreduce( + local_send_id_to_needed_local_recv_ids, dict_union_mpi_op) + finally: + dict_union_mpi_op.Free() + + return comm_to_needed_comms + +# }}} + + +# {{{ verify_distributed_dag_pre_partition + +def verify_distributed_dag_pre_partition(mpi_communicator: mpi4py.MPI.Comm, + outputs: DictOfNamedArrays) -> None: + """ + Verify that a global, unpartitioned graph does not contain a cycle. + + .. warning:: + + This is an MPI-collective operation. + """ + my_rank = mpi_communicator.rank + root_rank = 0 + + comm_to_needed_comms = _get_comm_to_needed_comms(mpi_communicator, outputs) + + if my_rank == root_rank: + from pytools.graph import compute_topological_order + compute_topological_order(comm_to_needed_comms) + + logger.info("verify_distributed_dag_pre_partition completed successfully.") + + +# }}} + + # {{{ verify_distributed_partition def verify_distributed_partition(mpi_communicator: mpi4py.MPI.Comm, diff --git a/test/test_distributed.py b/test/test_distributed.py index 95f220b7e..45664061a 100644 --- a/test/test_distributed.py +++ b/test/test_distributed.py @@ -415,6 +415,13 @@ def _do_verify_distributed_partition(ctx_factory): outputs = pt.make_dict_of_named_arrays({"out": send+send2}) distributed_parts = pt.find_distributed_partition(outputs) + if rank == 0: + from pytools.graph import CycleError + with pytest.raises(CycleError): + pt.verify_distributed_dag_pre_partition(comm, outputs) + else: + pt.verify_distributed_dag_pre_partition(comm, outputs) + if rank == 0: with pytest.raises(PartitionInducedCycleError): pt.verify_distributed_partition(comm, distributed_parts)