Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions examples/distributed.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
4 changes: 3 additions & 1 deletion pytato/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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",

Expand Down
104 changes: 102 additions & 2 deletions pytato/distributed/partition.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
Expand Down Expand Up @@ -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(
Expand Down
81 changes: 78 additions & 3 deletions pytato/distributed/verify.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -48,6 +50,7 @@

if TYPE_CHECKING:
import mpi4py.MPI
from mpi4py import MPI


# {{{ data structures
Expand Down Expand Up @@ -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,
Expand Down
7 changes: 7 additions & 0 deletions test/test_distributed.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down