diff --git a/tests/benchmark_multiple_processes_all_to_all.py b/tests/benchmark_multiple_processes_all_to_all.py new file mode 100644 index 000000000..aee6404f9 --- /dev/null +++ b/tests/benchmark_multiple_processes_all_to_all.py @@ -0,0 +1,240 @@ +import argparse +import multiprocessing + +from utils_all_to_all import ( + asyncio_process, + tornado_process, + ucx_process, + uvloop_process, +) + +from dask.utils import parse_bytes + +import ucp + + +def parse_args(): + parser = argparse.ArgumentParser(description="All-to-all benchmark") + parser.add_argument( + "--monitor", + default=False, + action="store_true", + help="Start a monitor process only. Requires --num-worker processes " + "started with --worker to connect to monitor. Default: disabled.", + ) + parser.add_argument( + "--worker", + default=False, + action="store_true", + help="Start a worker process only. Requires a --monitor process " + "to connect to, with an address specified with --monitor-address. " + "Default: disabled.", + ) + parser.add_argument( + "--enable-monitor", + default=False, + action="store_true", + help="Use a monitor process to synchronize workers when in single-node " + "mode (when neither --monitor or --worker are requested), otherwise " + "synchronize processes via multiprocessing shared memory. Default: " + "disabled.", + ) + parser.add_argument( + "--listen-interface", + default=None, + help="Interface where monitor (if --monitor), or worker (if --worker), " + "or all (in single-node mode, i.e., no --monitor or --worker are " + "specified) will listen for connections.", + ) + parser.add_argument( + "--monitor-address", + metavar="IP:PORT", + default=None, + help="Address where --monitor process is listening to in the HOST:PORT " + "format.", + ) + parser.add_argument( + "--port", + default=None, + type=int, + help="Port where --monitor will listen. Only applies to --monitor " + "process, --worker process should still use --monitor-address to " + "specify the monitor address. Default: random port.", + ) + parser.add_argument( + "--num-workers", + default=2, + type=int, + help="Number of workers to start in single-node mode, or number of " + "workers that --monitor process will wait to connect before starting " + "transfers between workers. Default: 2.", + ) + parser.add_argument( + "--endpoints-per-worker", + default=1, + type=int, + help="Number of simultaneous endpoints between each worker pair that " + "will send and receive data. In a case where --num-workers=2, this " + "translates to Worker1 creating two endpoints connecting to the " + "listener on Worker2, with Worker2 creating another two endpoints " + "connecting to the listener on Worker1, totalling four endpoint pairs " + "sending and receiving benchmark data simultaneously. Default: 1.", + ) + parser.add_argument( + "--communication-lib", + default="ucx", + type=str, + help="Communication library to benchmark. Options are " + "'ucx', 'asyncio', 'tornado', 'uvloop'. Default: 'ucx'.", + ) + parser.add_argument( + "--size", + metavar="BYTES", + default="1 MiB", + type=parse_bytes, + help="Size to be passed for data generation function. Default: '1 Mb'.", + ) + parser.add_argument( + "--iterations", + default=15, + type=int, + help="Number of iterations of data transfers per worker pair. " + "Each iteration consists in sending and receiving the same " + "data amount. Default: 15.", + ) + parser.add_argument( + "--gather-send-recv", + default=False, + action="store_true", + help="If disabled (default), send and receive operations will " + "be awaited individually, otherwise they are launched " + "simultaneously via an asyncio.gather or gen.multi operation. " + "Default: disabled.", + ) + + args = parser.parse_args() + if args.worker and args.monitor: + raise RuntimeError("--monitor and --worker can't be defined together.") + return args + + +def main(): + args = parse_args() + + num_workers = args.num_workers + endpoints_per_worker = args.endpoints_per_worker + listener_address = ucp.get_address(ifname=args.listen_interface) + if args.monitor_address is not None: + monitor_address, monitor_port = args.monitor_address.split(":") + monitor_port = int(monitor_port) + + if args.communication_lib == "ucx": + communication_func = ucx_process + elif args.communication_lib == "asyncio": + communication_func = asyncio_process + elif args.communication_lib == "uvloop": + communication_func = uvloop_process + elif args.communication_lib == "tornado": + communication_func = tornado_process + else: + raise ValueError( + f"Communication library {args.communication_lib} not supported" + ) + + if args.monitor is True: + communication_func( + listener_address, + num_workers, + endpoints_per_worker, + True, + args.port, + args.size, + args.iterations, + args.gather_send_recv, + shm_sync=False, + ) + elif args.worker is True: + communication_func( + listener_address, + num_workers, + endpoints_per_worker, + False, + monitor_port, + args.size, + args.iterations, + args.gather_send_recv, + shm_sync=False, + ) + else: + ctx = multiprocessing.get_context("spawn") + + signal = ctx.Array("i", [0, 0]) + ports = ctx.Array("i", range(num_workers)) + lock = ctx.Lock() + + monitor_port = 0 + + if args.enable_monitor: + monitor_process = ctx.Process( + name="worker", + target=communication_func, + args=[ + listener_address, + num_workers, + endpoints_per_worker, + True, + 0, + args.size, + args.iterations, + args.gather_send_recv, + ], + kwargs={ + "shm_sync": True, + "signal": signal, + "ports": ports, + "lock": lock, + }, + ) + monitor_process.start() + + while signal[0] == 0: + pass + + monitor_port = signal[0] + + worker_processes = [] + for worker_num in range(num_workers): + worker_process = ctx.Process( + name="worker", + target=communication_func, + args=[ + listener_address, + num_workers, + endpoints_per_worker, + False, + monitor_port, + args.size, + args.iterations, + args.gather_send_recv, + ], + kwargs={ + "shm_sync": not args.enable_monitor, + "signal": signal, + "ports": ports, + "lock": lock, + }, + ) + worker_process.start() + worker_processes.append(worker_process) + + for worker_process in worker_processes: + worker_process.join() + + if args.enable_monitor: + monitor_process.join() + + assert worker_process.exitcode == 0 + + +if __name__ == "__main__": + main() diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py new file mode 100644 index 000000000..6f139a7fc --- /dev/null +++ b/tests/test_multiple_processes_all_to_all.py @@ -0,0 +1,104 @@ +import multiprocessing + +import pytest +from utils_all_to_all import ( + asyncio_process, + tornado_process, + ucx_process, + uvloop_process, +) + +import ucp + + +def _test_send_recv_cu( + num_workers, endpoints_per_worker, enable_monitor, size, iterations, communication +): + ctx = multiprocessing.get_context("spawn") + + listener_address = ucp.get_address() + monitor_port = 0 + + signal = ctx.Array("i", [0, 0]) + ports = ctx.Array("i", range(num_workers)) + lock = ctx.Lock() + + if enable_monitor: + monitor_process = ctx.Process( + name="worker", + target=communication, + args=[ + listener_address, + num_workers, + endpoints_per_worker, + True, + 0, + size, + iterations, + False, + ], + kwargs={"shm_sync": True, "signal": signal, "ports": ports, "lock": lock}, + ) + monitor_process.start() + + while signal[0] == 0: + pass + + monitor_port = signal[0] + + worker_processes = [] + for worker_num in range(num_workers): + worker_process = ctx.Process( + name="worker", + target=communication, + args=[ + listener_address, + num_workers, + endpoints_per_worker, + False, + monitor_port, + size, + iterations, + False, + ], + kwargs={ + "shm_sync": not enable_monitor, + "signal": signal, + "ports": ports, + "lock": lock, + }, + ) + worker_process.start() + worker_processes.append(worker_process) + + for worker_process in worker_processes: + worker_process.join() + + if enable_monitor: + monitor_process.join() + + assert worker_process.exitcode == 0 + + +@pytest.mark.parametrize("num_workers", [2, 4, 8]) +@pytest.mark.parametrize("endpoints_per_worker", [1]) +@pytest.mark.parametrize("enable_monitor", [True, False]) +@pytest.mark.parametrize("size", [2 ** 20]) +@pytest.mark.parametrize("iterations", [5]) +@pytest.mark.parametrize( + "communication", [ucx_process, asyncio_process, uvloop_process, tornado_process] +) +def test_send_recv_cu( + num_workers, endpoints_per_worker, enable_monitor, size, iterations, communication +): + if communication == uvloop_process: + pytest.importorskip("uvloop", reason="uvloop not installed") + + _test_send_recv_cu( + num_workers, + endpoints_per_worker, + enable_monitor, + size, + iterations, + communication, + ) diff --git a/tests/utils_all_to_all.py b/tests/utils_all_to_all.py new file mode 100644 index 000000000..892c70704 --- /dev/null +++ b/tests/utils_all_to_all.py @@ -0,0 +1,528 @@ +import asyncio +from time import monotonic + +import numpy as np +from tornado import gen +from tornado.ioloop import IOLoop +from utils_comm_libs import ( + AsyncioCommConnection, + AsyncioCommServer, + TornadoTCPConnection, + TornadoTCPServer, + UCXConnection, + UCXServer, +) + +from dask.utils import format_bytes +from distributed.utils import nbytes + +import ucp + +OP_NONE = 0 +OP_WORKER_LISTENING = 1 +OP_CLUSTER_READY = 2 +OP_WORKER_COMPLETED = 3 +OP_SHUTDOWN = 4 + + +class UCXProcess: + def __init__( + self, + listener_address, + num_workers, + endpoints_per_worker, + monitor_port, + size, + iterations, + gather_send_recv, + transfer_to_cache, + shm_sync=True, + signal=None, + ports=None, + lock=None, + ): + self.listener_address = listener_address + self.num_workers = num_workers + self.endpoints_per_worker = endpoints_per_worker + self.monitor_port = monitor_port + + # UCX only effectively creates endpoints at first transfer, but this + # isn't necessary for tornado/asyncio. + self.transfer_to_cache = transfer_to_cache + + self.size = size + self.iterations = iterations + self.gather_send_recv = gather_send_recv + + self.shm_sync = shm_sync + self.signal = signal + self.ports = ports + self.lock = lock + + self.conns = dict() + self.connections = set() + self.bytes_bandwidth = dict() + + self.cluster_started = False + + async def _sleep(self, delay): + await asyncio.sleep(delay) + + async def _gather(self, tasks): + await asyncio.gather(*tasks) + + async def _transfer(self, ep, msg2send, send_first=True): + time_per_iteration = [] + for i in range(self.iterations): + t = monotonic() + if self.gather_send_recv: + msgs = [ep.recv(), ep.send(msg2send)] + await self._gather(msgs) + else: + # This seems to be faster! + if send_first: + await ep.send(msg2send) + await ep.recv() + else: + await ep.recv() + await ep.send(msg2send) + time_per_iteration.append(monotonic() - t) + return time_per_iteration + + async def _listener(self, ep): + message = np.arange(self.size, dtype=np.uint8) + + await self._transfer(ep, message) + + async def _client(self, my_port, worker_address, ep, cache_only=False): + message = np.arange(self.size, dtype=np.uint8) + send_recv_bytes = nbytes(message) * 2 + + time_per_iteration = await self._transfer(ep, message, send_first=False) + + if cache_only is False: + self.bytes_bandwidth[worker_address] += list( + (send_recv_bytes, send_recv_bytes / t) for t in time_per_iteration + ) + + def _init(self): + ucp.init() + + async def _create_listener(self, host, port=None, cb=None): + return await UCXServer.start_server(host, port, cb or self._listener) + + async def _monitor_listener_cb(self, ep): + return None + + async def _create_endpoint(self, host, port): + my_port = self.listener.port + my_task = "Worker" + remote_task = "Worker" + + while True: + try: + ep = await UCXConnection.open_connection(host, port) + return ep + except ucp.exceptions.UCXCanceled as e: + print( + "%s[%d]->%s[%d] Failed: %s" + % (my_task, my_port, remote_task, port, e), + flush=True, + ) + await self._sleep(0.1) + + def get_connections(self): + return self.listener._connections + + async def _wait_for_workers(self): + if self.shm_sync: + with self.lock: + self.ports[self.signal[0]] = self.listener.port + self.signal[0] += 1 + + while self.signal[0] != self.num_workers: + await self._sleep(0.1) + else: + self.monitor_ep = await self._create_endpoint( + self.listener_address, self.monitor_port + ) + await self.monitor_ep.send( + (OP_WORKER_LISTENING, self.listener_address, self.listener.port) + ) + _, worker_addresses = await self.monitor_ep.recv() + assert worker_addresses["data"][0] == OP_CLUSTER_READY + self.worker_addresses = worker_addresses["data"][1] + if len(self.worker_addresses) != self.num_workers: + raise ValueError( + f"Wrong number of workers, {self.num_workers} expected by this " + f"worker, but monitor reported {len(self.worker_addresses)}. " + "Make sure monitor and all workers specify the same number of " + "workers." + ) + + async def _wait_for_completion(self): + if self.shm_sync: + with self.lock: + self.ports[self.signal[1]] = self.listener.port + self.signal[1] += 1 + + while self.signal[1] != self.num_workers: + await self._sleep(0) + else: + await self.monitor_ep.send((OP_WORKER_COMPLETED,)) + _, shutdown = await self.monitor_ep.recv() + shutdown = shutdown["data"] + assert shutdown[0] == OP_SHUTDOWN + + async def _create_all_endpoints(self): + for i in range(self.endpoints_per_worker): + client_tasks = [] + # Create endpoints to all other workers + if self.shm_sync: + for remote_port in list(self.ports): + if remote_port == self.listener.port: + continue + + remote_address = (self.listener_address, remote_port) + + ep = await self._create_endpoint(*remote_address) + self.bytes_bandwidth[remote_address] = [] + self.conns[(remote_address, i)] = ep + + if self.transfer_to_cache: + client_tasks.append( + self._client( + self.listener.port, remote_address, ep, cache_only=True + ) + ) + else: + for worker_address in self.worker_addresses: + if ( + worker_address[0] == self.listener_address + and worker_address[1] == self.listener.port + ): + continue + + ep = await self._create_endpoint(*worker_address) + self.bytes_bandwidth[worker_address] = [] + self.conns[(worker_address, i)] = ep + + if self.transfer_to_cache: + client_tasks.append( + self._client( + self.listener.port, worker_address, ep, cache_only=True + ) + ) + + if self.transfer_to_cache: + await self._gather(client_tasks) + + async def _wait_for_connections_cache(self): + # Wait until listener->ep connections have all been cached + while len(self.get_connections()) != self.endpoints_per_worker * ( + self.num_workers - 1 + ): + await self._sleep(0.1) + + async def _exchange_messages(self): + # Exchange messages with other workers + client_tasks = [] + listener_tasks = [] + for (worker_address, _), ep in self.conns.items(): + client_tasks.append(self._client(self.listener.port, worker_address, ep)) + for listener_ep in self.get_connections(): + listener_tasks.append(self._listener(listener_ep)) + all_tasks = client_tasks + listener_tasks + await self._gather(all_tasks) + + async def _close_connections_and_listener(self): + for conn in self.get_connections(): + await conn.close() + + self.listener.close() + + # Wait for a shutdown signal from monitor + try: + while not self.listener.closed(): + await self._sleep(0.1) + except ucp.UCXCloseError: + pass + + async def run_worker(self): + self._init() + + # Start listener + self.listener = await self._create_listener(self.listener_address) + + await self._wait_for_workers() + + await self._create_all_endpoints() + + await self._wait_for_connections_cache() + + await self._exchange_messages() + + await self._wait_for_completion() + + await self._close_connections_and_listener() + + async def run_monitor(self): + self._init() + + self.listener = await self._create_listener( + self.listener_address, self.monitor_port, self._monitor_listener_cb, + ) + + if self.shm_sync: + with self.lock: + self.signal[0] = self.listener.port + else: + print(f"Monitor listening at {self.listener_address}:{self.listener.port}") + + # Wait for all workers to connect + while len(self.get_connections()) != self.num_workers: + if not self.shm_sync: + print( + "Waiting for all workers to connect, " + f"{len(self.get_connections())} of " + f"{self.num_workers} workers connected." + ) + await self._sleep(1) + + # Get all worker addresses + worker_addresses = [] + for conn in self.get_connections(): + _, address = await conn.recv() + address = address["data"] + assert address[0] == OP_WORKER_LISTENING + worker_addresses.append((address[1], address[2])) + + # Send a list of all worker addresses to each worker, indicating the cluster + # is ready + for conn in self.get_connections(): + await conn.send([OP_CLUSTER_READY, worker_addresses]) + + # Wait for all workers to complete + for conn in self.get_connections(): + _, complete = await conn.recv() + complete = complete["data"] + assert int(complete[0]) == OP_WORKER_COMPLETED + + # Signal all workers to shutdown + for conn in self.get_connections(): + await conn.send((OP_SHUTDOWN,)) + + await self._close_connections_and_listener() + + def get_results(self): + for remote_address, bb in self.bytes_bandwidth.items(): + local_address = (self.listener_address, self.listener.port) + total_bytes = sum(b[0] for b in bb) + avg_bandwidth = np.mean(list(b[1] for b in bb)) + median_bandwidth = np.median(list(b[1] for b in bb)) + print( + "[%s -> %s] Transferred bytes: %s, average bandwidth: %s/s, " + "median bandwidth: %s/s" + % ( + ":".join([str(i) for i in local_address]), + ":".join([str(i) for i in remote_address]), + format_bytes(total_bytes), + format_bytes(avg_bandwidth), + format_bytes(median_bandwidth), + ) + ) + + +class TornadoProcess(UCXProcess): + def __init__( + self, + listener_address, + num_workers, + endpoints_per_worker, + monitor_port, + size, + iterations, + gather_send_recv, + shm_sync=True, + signal=None, + ports=None, + lock=None, + ): + super().__init__( + listener_address=listener_address, + num_workers=num_workers, + endpoints_per_worker=endpoints_per_worker, + monitor_port=monitor_port, + size=size, + iterations=iterations, + gather_send_recv=gather_send_recv, + transfer_to_cache=False, + shm_sync=shm_sync, + signal=signal, + ports=ports, + lock=lock, + ) + + async def _sleep(self, delay): + await gen.sleep(delay) + + async def _gather(self, tasks): + await gen.multi(tasks) + + def _init(self): + return + + async def _create_listener(self, host, port=None, cb=None): + return await TornadoTCPServer.start_server(host, port=port) + + async def _create_endpoint(self, host, port): + return await TornadoTCPConnection.connect(host, port) + + +class AsyncioProcess(UCXProcess): + def __init__( + self, + listener_address, + num_workers, + endpoints_per_worker, + monitor_port, + size, + iterations, + gather_send_recv, + shm_sync=True, + signal=None, + ports=None, + lock=None, + ): + super().__init__( + listener_address=listener_address, + num_workers=num_workers, + endpoints_per_worker=endpoints_per_worker, + monitor_port=monitor_port, + size=size, + iterations=iterations, + gather_send_recv=gather_send_recv, + transfer_to_cache=False, + shm_sync=shm_sync, + signal=signal, + ports=ports, + lock=lock, + ) + + def _init(self): + return + + async def _create_listener(self, host, port=None, cb=None): + return await AsyncioCommServer.start_server(host, port) + + async def _create_endpoint(self, host, port): + return await AsyncioCommConnection.open_connection(host, port) + + +def ucx_process( + listener_address, + num_workers, + endpoints_per_worker, + is_monitor, + monitor_port, + size, + iterations, + gather_send_recv, + shm_sync=True, + signal=None, + ports=None, + lock=None, +): + w = UCXProcess( + listener_address, + num_workers, + endpoints_per_worker, + monitor_port, + size=size, + iterations=iterations, + gather_send_recv=gather_send_recv, + transfer_to_cache=True, + shm_sync=shm_sync, + signal=signal, + ports=ports, + lock=lock, + ) + run_func = w.run_monitor if is_monitor else w.run_worker + asyncio.get_event_loop().run_until_complete(run_func()) + w.get_results() + + +def asyncio_process( + listener_address, + num_workers, + endpoints_per_worker, + is_monitor, + monitor_port, + size, + iterations, + gather_send_recv, + shm_sync=True, + signal=None, + ports=None, + lock=None, + loop=None, +): + w = AsyncioProcess( + listener_address, + num_workers, + endpoints_per_worker, + monitor_port, + size=size, + iterations=iterations, + gather_send_recv=gather_send_recv, + shm_sync=shm_sync, + signal=signal, + ports=ports, + lock=lock, + ) + run_func = w.run_monitor if is_monitor else w.run_worker + if loop is None: + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + loop.run_until_complete(run_func()) + w.get_results() + + +def uvloop_process( + *args, **kwargs, +): + import uvloop + + return asyncio_process(*args, **kwargs, loop=uvloop.new_event_loop()) + + +def tornado_process( + listener_address, + num_workers, + endpoints_per_worker, + is_monitor, + monitor_port, + size, + iterations, + gather_send_recv, + shm_sync=True, + signal=None, + ports=None, + lock=None, +): + w = TornadoProcess( + listener_address, + num_workers, + endpoints_per_worker, + monitor_port, + size=size, + iterations=iterations, + gather_send_recv=gather_send_recv, + shm_sync=shm_sync, + signal=signal, + ports=ports, + lock=lock, + ) + run_func = w.run_monitor if is_monitor else w.run_worker + IOLoop.current().run_sync(run_func) + w.get_results() diff --git a/tests/utils_comm_libs.py b/tests/utils_comm_libs.py new file mode 100644 index 000000000..b7dbdd0a8 --- /dev/null +++ b/tests/utils_comm_libs.py @@ -0,0 +1,322 @@ +import asyncio +import struct + +import numpy as np +from tornado.iostream import StreamClosedError +from tornado.tcpclient import TCPClient +from tornado.tcpserver import TCPServer + +from distributed.comm.utils import from_frames, to_frames +from distributed.protocol import to_serialize +from distributed.protocol.utils import pack_frames_prelude, unpack_frames +from distributed.utils import nbytes + +import ucp + + +class TornadoTCPConnection: + def __init__(self, stream, client=None): + self._client = client + self.stream = stream + self._closed = False + + @classmethod + async def connect(cls, host, port): + client = TCPClient() + stream = await client.connect(host, port, max_buffer_size=2 ** 30) + stream.set_nodelay(True) + return cls(stream, client=client) + + async def send(self, message): + stream = self.stream + if stream is None: + raise StreamClosedError() + + msg = {"data": to_serialize(message)} + frames = await to_frames(msg, serializers=("cuda", "dask", "pickle")) + + frames_nbytes = [nbytes(f) for f in frames] + frames_nbytes_total = sum(frames_nbytes) + + header = pack_frames_prelude(frames) + header = struct.pack("Q", nbytes(header) + frames_nbytes_total) + header + + frames = [header, *frames] + frames_nbytes = [nbytes(header), *frames_nbytes] + frames_nbytes_total += frames_nbytes[0] + + if frames_nbytes_total < 2 ** 17: + frames = [b"".join(frames)] + frames_nbytes = [frames_nbytes_total] + + try: + for each_frame_nbytes, each_frame in zip(frames_nbytes, frames): + if each_frame_nbytes: + if stream._write_buffer is None: + raise StreamClosedError() + + if isinstance(each_frame, memoryview): + each_frame = memoryview(each_frame).cast("B") + + stream._write_buffer.append(each_frame) + stream._total_write_index += each_frame_nbytes + + stream.write(b"") + except StreamClosedError: + self.stream = None + self._closed = True + except Exception() as e: + raise e + + return frames_nbytes_total + + async def recv(self): + stream = self.stream + if stream is None: + raise Exception("Connection closed") + + fmt = "Q" + fmt_size = struct.calcsize(fmt) + + try: + frames_nbytes = await stream.read_bytes(fmt_size) + (frames_nbytes,) = struct.unpack(fmt, frames_nbytes) + + frames = bytearray(frames_nbytes) + n = await stream.read_into(frames) + assert n == frames_nbytes, (n, frames_nbytes) + except StreamClosedError: + self.stream = None + self._closed = True + except Exception as e: + raise e + else: + try: + frames = unpack_frames(frames) + + msg = await from_frames( + frames, + deserializers=("cuda", "dask", "pickle", "error"), + allow_offload=True, + ) + except EOFError: + raise Exception("aborted stream on truncated data") + return frames, msg + + async def close(self): + self.stream.close() + self._closed = True + + def closed(self): + return self._closed + + +class TornadoTCPServer: + def __init__(self, server, connections, port): + server.handle_stream = self._handle_stream + self.server = server + self._connections = connections + self._port = port + + async def _handle_stream(self, stream, address): + self._connections.append(TornadoTCPConnection(stream)) + + @classmethod + async def start_server(cls, host, port): + connections = [] + + server = TCPServer(max_buffer_size=2 ** 30) + + if port is None or port == 0: + + def _try_listen(server, host): + while True: + try: + import random + + port = random.randint(10000, 60000) + server.listen(port, host) + return port + except OSError: + pass + + port = _try_listen(server, host) + else: + server.listen(port, host) + + server.start() + + return cls(server, connections, port) + + def get_connections(self): + return self._connections + + @property + def port(self): + return self._port + + def close(self): + self.server.stop() + self._closed = True + + def closed(self): + return self._closed + + +class AsyncioCommConnection: + def __init__(self, reader, writer): + self.reader = reader + self.writer = writer + + @classmethod + async def open_connection(cls, host, port): + reader, writer = await asyncio.open_connection(host, port, limit=2 ** 30) + return cls(reader, writer) + + async def send(self, message): + msg = {"data": to_serialize(message)} + frames = await to_frames(msg, serializers=("cuda", "dask", "pickle")) + + nframes = len(frames) + self.writer.write(struct.pack("Q", nframes)) + sizes = list(nbytes(f) for f in frames) + self.writer.write(struct.pack(nframes * "Q", *sizes)) + for f in frames: + self.writer.write(f) + await self.writer.drain() + + async def recv(self): + nframes = await self.reader.readexactly(struct.calcsize("Q")) + nframes = struct.unpack("Q", nframes) + sizes = await self.reader.readexactly(struct.calcsize(nframes[0] * "Q")) + sizes = struct.unpack(nframes[0] * "Q", sizes) + frames = [] + for size in sizes: + frames.append(await self.reader.readexactly(size)) + + msg = await from_frames( + frames, + deserializers=("cuda", "dask", "pickle", "error"), + allow_offload=True, + ) + return frames, msg + + async def close(self): + self.writer.close() + + def closed(self): + return self.writer.is_closing() + + +class AsyncioCommServer: + def __init__(self, server, connections): + self.server = server + self._connections = connections + self._port = self.server.sockets[0].getsockname()[1] + + @classmethod + async def start_server(cls, host, port): + connections = [] + + def _server_callback(reader, writer): + connections.append(AsyncioCommConnection(reader, writer)) + + server = await asyncio.start_server( + _server_callback, host, port, limit=2 ** 30, + ) + return cls(server, connections) + + def get_connections(self): + return self._connections + + @property + def port(self): + return self._port + + def close(self): + self.server.close() + + def closed(self): + return not self.server.is_serving() + + +class UCXConnection: + def __init__(self, ep): + self.ep = ep + + @classmethod + async def open_connection(cls, host, port): + ep = await ucp.create_endpoint(host, port) + return cls(ep) + + async def send(self, message): + msg = {"data": to_serialize(message)} + frames = await to_frames(msg, serializers=("cuda", "dask", "pickle")) + + nframes = len(frames) + await self.ep.send(struct.pack("Q", nframes)) + + sizes = list(nbytes(f) for f in frames) + await self.ep.send(struct.pack(nframes * "Q", *sizes)) + + for f in frames: + await self.ep.send(f) + + async def recv(self): + nframes = np.empty((struct.calcsize("Q"),), dtype="u1") + await self.ep.recv(nframes) + nframes = struct.unpack("Q", nframes) + + sizes = np.empty((struct.calcsize(nframes[0] * "Q"),), dtype="u1") + await self.ep.recv(sizes) + sizes = struct.unpack(nframes[0] * "Q", sizes) + + frames = [] + for size in sizes: + frame = np.empty((size,), dtype="u1") + await self.ep.recv(frame) + frames.append(frame) + + msg = await from_frames( + frames, + deserializers=("cuda", "dask", "pickle", "error"), + allow_offload=True, + ) + return frames, msg + + async def close(self): + await self.ep.close() + + def closed(self): + return self.ep.closed() + + +class UCXServer: + def __init__(self, server, connections): + self.server = server + self._connections = connections + + @classmethod + async def start_server(cls, host, port, listener_func): + connections = [] + + async def _server_callback(ep): + conn = UCXConnection(ep) + connections.append(conn) + await listener_func(conn) + + server = ucp.create_listener(_server_callback, port) + return cls(server, connections) + + def get_connections(self): + return self._connections + + @property + def port(self): + return self.server.port + + def close(self): + self.server.close() + + def closed(self): + return self.server.closed()