From 3c0a12afa39662e0ed748f50b4f8899d7d250659 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Mon, 8 Mar 2021 15:12:16 -0800 Subject: [PATCH 01/42] Add multiprocess test to perform all-to-all ep creation --- tests/test_multiple_processes_all_to_all.py | 210 ++++++++++++++++++++ 1 file changed, 210 insertions(+) create mode 100644 tests/test_multiple_processes_all_to_all.py 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..e3fc103b7 --- /dev/null +++ b/tests/test_multiple_processes_all_to_all.py @@ -0,0 +1,210 @@ +import asyncio +import multiprocessing +import random +import sys + +import numpy as np +import pytest + +import ucp + +OP_BYTES = 1 +PORT_BYTES = 2 + +OP_NONE = 0 +OP_WORKER_LISTENING = 1 +OP_WORKER_COMPLETED = 2 +OP_CLUSTER_READY = 3 +OP_SHUTDOWN = 4 + + +def generate_op_message(op, port): + op_bytes = op.to_bytes(OP_BYTES, sys.byteorder) + port_bytes = port.to_bytes(PORT_BYTES, sys.byteorder) + return bytearray(b"".join([op_bytes, port_bytes])) + + +def parse_op_message(msg): + op = int.from_bytes(msg[0:OP_BYTES], sys.byteorder) + port = int.from_bytes(msg[OP_BYTES : OP_BYTES + PORT_BYTES], sys.byteorder) + return {"op": op, "port": port} + + +def worker(my_port, monitor_port, all_ports): + ucp.init() + + global cluster_started + cluster_started = False + + async def _worker(my_port, all_ports): + def _register_cluster_started(): + global cluster_started + cluster_started = True + + async def _listener(ep): + op_msg = generate_op_message(OP_NONE, 0) + msg2send = np.arange(10) + msg2recv = np.empty_like(msg2send) + + msgs = [ep.recv(op_msg), ep.send(msg2send), ep.recv(msg2recv)] + await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) + + op = parse_op_message(op_msg)["op"] + + if op == OP_SHUTDOWN: + await ep.close() + listener.close() + if op == OP_CLUSTER_READY: + _register_cluster_started() + + async def _client(port): + op_msg = generate_op_message(OP_NONE, 0) + msg2send = np.arange(10) + msg2recv = np.empty_like(msg2send) + + ep = await ucp.create_endpoint(ucp.get_address(), port) + msgs = [ep.send(op_msg), ep.recv(msg2send), ep.send(msg2recv)] + await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) + + await asyncio.sleep(2) + + async def _signal_monitor(monitor_port, my_port, op): + op_msg = generate_op_message(op, my_port) + ack_msg = bytearray(2) + + ep = await ucp.create_endpoint(ucp.get_address(), monitor_port) + msgs = [ep.send(op_msg), ep.recv(ack_msg)] + await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) + + # Start listener + listener = ucp.create_listener(_listener, port=my_port) + + # Signal monitor that worker is listening + await _signal_monitor(monitor_port, my_port, op=OP_WORKER_LISTENING) + + while not cluster_started: + await asyncio.sleep(0.1) + + # Create endpoints to all other workers + clients = [] + for port in all_ports: + clients.append(_client(port)) + await asyncio.gather(*clients, loop=asyncio.get_event_loop()) + + # Signal monitor that worker is completed + await _signal_monitor(monitor_port, my_port, op=OP_WORKER_COMPLETED) + + # Wait for a shutdown signal from monitor + try: + while not listener.closed(): + await asyncio.sleep(0.1) + except ucp.UCXCloseError: + pass + + asyncio.get_event_loop().run_until_complete(_worker(my_port, all_ports)) + + +def monitor(monitor_port, worker_ports): + ucp.init() + + listening_worker_ports = [] + completed_worker_ports = [] + + async def _monitor(monitor_port, worker_ports): + def _register(op, port): + if op == OP_WORKER_LISTENING: + listening_worker_ports.append(port) + elif op == OP_WORKER_COMPLETED: + completed_worker_ports.append(port) + + async def _listener(ep): + op_msg = generate_op_message(OP_NONE, 0) + ack_msg = bytearray(int(888).to_bytes(2, sys.byteorder)) + + # Sending an ack_msg prevents the other ep from closing too + # early, ultimately leading this process to hang. + msgs = [ep.recv(op_msg), ep.send(ack_msg)] + await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) + + # worker_op == 0 for started, worker_op == 1 for completed + op_msg = parse_op_message(op_msg) + worker_op = op_msg["op"] + worker_port = op_msg["port"] + + _register(worker_op, worker_port) + + async def _send_op(op, port): + op_msg = generate_op_message(op, port) + msg2send = np.arange(10) + msg2recv = np.empty_like(msg2send) + + ep = await ucp.create_endpoint(ucp.get_address(), port) + msgs = [ep.send(op_msg), ep.send(msg2send), ep.recv(msg2recv)] + await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) + + # Start monitor's listener + listener = ucp.create_listener(_listener, port=monitor_port) + + # Wait until all workers signal they are listening + while len(listening_worker_ports) != len(worker_ports): + await asyncio.sleep(0.1) + + # Send cluster ready message to all workers + ready_signals = [] + for port in listening_worker_ports: + ready_signals.append(_send_op(OP_CLUSTER_READY, port)) + await asyncio.gather(*ready_signals, loop=asyncio.get_event_loop()) + + # Wait until all workers signal completion + while len(completed_worker_ports) != len(worker_ports): + await asyncio.sleep(0.1) + + # Send shutdown message to all workers + close = [] + for port in completed_worker_ports: + close.append(_send_op(OP_SHUTDOWN, port)) + await asyncio.gather(*close, loop=asyncio.get_event_loop()) + + listener.close() + + asyncio.get_event_loop().run_until_complete(_monitor(monitor_port, worker_ports)) + + +@pytest.mark.parametrize("num_workers", [1, 2, 4, 8]) +def test_send_recv_cu(num_workers): + # One additional port for monitor + num_ports = num_workers + 1 + + ports = set() + while len(ports) != num_ports: + missing_ports = num_ports - len(ports) + ports = ports.union( + [random.randint(13000, 23000) for n in range(missing_ports)] + ) + ports = list(ports) + + monitor_port = ports[0] + worker_ports = ports[1:] + + ctx = multiprocessing.get_context("spawn") + + monitor_process = ctx.Process( + name="monitor", target=monitor, args=[monitor_port, worker_ports] + ) + monitor_process.start() + + worker_processes = [] + for port in worker_ports: + worker_process = ctx.Process( + name="worker", target=worker, args=[port, monitor_port, worker_ports] + ) + worker_process.start() + worker_processes.append(worker_process) + + for worker_process in worker_processes: + worker_process.join() + + monitor_process.join() + + assert worker_process.exitcode == 0 + assert monitor_process.exitcode == 0 From 35664da55fb0d26c1381425f31f9bdaf6fa1bb4b Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Wed, 10 Mar 2021 07:52:42 -0800 Subject: [PATCH 02/42] Add support for persistent endpoints in multiprocess all_to_all test --- tests/test_multiple_processes_all_to_all.py | 152 ++++++++++++++++---- 1 file changed, 128 insertions(+), 24 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index e3fc103b7..e5618fb13 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -17,6 +17,8 @@ OP_CLUSTER_READY = 3 OP_SHUTDOWN = 4 +PersistentEndpoints = False + def generate_op_message(op, port): op_bytes = op.to_bytes(OP_BYTES, sys.byteorder) @@ -30,9 +32,28 @@ def parse_op_message(msg): return {"op": op, "port": port} +async def create_endpoint_retry(my_port, remote_port, my_task, remote_task): + while True: + try: + ep = await ucp.create_endpoint(ucp.get_address(), remote_port) + return ep + except ucp.exceptions.UCXCanceled as e: + print( + "%s[%d]->%s[%d] Failed: %s" + % (my_task, my_port, remote_task, remote_port, e), + flush=True, + ) + await asyncio.sleep(0.1) + + def worker(my_port, monitor_port, all_ports): ucp.init() + listener_eps = [] + + global listener_monitor_ep + listener_monitor_ep = None + global cluster_started cluster_started = False @@ -41,7 +62,11 @@ def _register_cluster_started(): global cluster_started cluster_started = True - async def _listener(ep): + async def _close_endpoints(): + for ep in listener_eps: + await ep.close() + + async def _listener(ep, cache_ep=False): op_msg = generate_op_message(OP_NONE, 0) msg2send = np.arange(10) msg2recv = np.empty_like(msg2send) @@ -51,48 +76,97 @@ async def _listener(ep): op = parse_op_message(op_msg)["op"] + if cache_ep and PersistentEndpoints: + if op == OP_NONE: + listener_eps.append(ep) + else: + global listener_monitor_ep + listener_monitor_ep = ep + if op == OP_SHUTDOWN: await ep.close() listener.close() if op == OP_CLUSTER_READY: _register_cluster_started() - async def _client(port): + async def _listener_cb(ep): + await _listener(ep, cache_ep=True) + + async def _client(port, ep=None): op_msg = generate_op_message(OP_NONE, 0) msg2send = np.arange(10) msg2recv = np.empty_like(msg2send) - ep = await ucp.create_endpoint(ucp.get_address(), port) - msgs = [ep.send(op_msg), ep.recv(msg2send), ep.send(msg2recv)] + if ep is None: + ep = await create_endpoint_retry(my_port, port, "Worker", "Worker") + msgs = [ep.send(op_msg), ep.recv(msg2recv), ep.send(msg2send)] await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) - await asyncio.sleep(2) - - async def _signal_monitor(monitor_port, my_port, op): + async def _signal_monitor(monitor_port, my_port, op, ep=None): op_msg = generate_op_message(op, my_port) ack_msg = bytearray(2) - ep = await ucp.create_endpoint(ucp.get_address(), monitor_port) + if ep is None: + #ep = await ucp.create_endpoint(ucp.get_address(), monitor_port) + ep = await create_endpoint_retry(my_port, monitor_port, "Worker", "Monitor") msgs = [ep.send(op_msg), ep.recv(ack_msg)] await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) # Start listener - listener = ucp.create_listener(_listener, port=my_port) + listener = ucp.create_listener(_listener_cb, port=my_port) # Signal monitor that worker is listening - await _signal_monitor(monitor_port, my_port, op=OP_WORKER_LISTENING) + monitor_ep = None + if PersistentEndpoints: + monitor_ep = await create_endpoint_retry( + my_port, monitor_port, "Worker", "Monitor" + ) + await _signal_monitor( + monitor_port, my_port, op=OP_WORKER_LISTENING, ep=monitor_ep + ) + else: + await _signal_monitor(monitor_port, my_port, op=OP_WORKER_LISTENING) while not cluster_started: await asyncio.sleep(0.1) - # Create endpoints to all other workers - clients = [] - for port in all_ports: - clients.append(_client(port)) - await asyncio.gather(*clients, loop=asyncio.get_event_loop()) + eps = [] + if PersistentEndpoints: + client_tasks = [] + # Create endpoints to all other workers + for remote_port in all_ports: + if remote_port == my_port: + continue + ep = await create_endpoint_retry( + my_port, remote_port, "Worker", "Worker" + ) + eps.append(ep) + client_tasks.append(_client(remote_port, ep)) + await asyncio.gather(*client_tasks, loop=asyncio.get_event_loop()) + + # Wait until listener_eps have all been cached + while len(listener_eps) != len(all_ports) - 1: + await asyncio.sleep(0.1) + else: + # Create endpoints to all other workers + client_tasks = [] + for port in all_ports: + if port == my_port: + continue + client_tasks.append(_client(port)) + await asyncio.gather(*client_tasks, loop=asyncio.get_event_loop()) # Signal monitor that worker is completed - await _signal_monitor(monitor_port, my_port, op=OP_WORKER_COMPLETED) + if PersistentEndpoints: + await _signal_monitor( + monitor_port, my_port, op=OP_WORKER_COMPLETED, ep=monitor_ep + ) + else: + await _signal_monitor(monitor_port, my_port, op=OP_WORKER_COMPLETED) + + # Wait for closing signal + if PersistentEndpoints: + await _listener(listener_monitor_ep) # Wait for a shutdown signal from monitor try: @@ -107,6 +181,7 @@ async def _signal_monitor(monitor_port, my_port, op): def monitor(monitor_port, worker_ports): ucp.init() + listener_eps = [] listening_worker_ports = [] completed_worker_ports = [] @@ -117,7 +192,10 @@ def _register(op, port): elif op == OP_WORKER_COMPLETED: completed_worker_ports.append(port) - async def _listener(ep): + async def _listener(ep, cache_ep=True): + if cache_ep and PersistentEndpoints: + listener_eps.append(ep) + op_msg = generate_op_message(OP_NONE, 0) ack_msg = bytearray(int(888).to_bytes(2, sys.byteorder)) @@ -126,35 +204,58 @@ async def _listener(ep): msgs = [ep.recv(op_msg), ep.send(ack_msg)] await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) - # worker_op == 0 for started, worker_op == 1 for completed op_msg = parse_op_message(op_msg) worker_op = op_msg["op"] worker_port = op_msg["port"] _register(worker_op, worker_port) - async def _send_op(op, port): + async def _listener_cb(ep): + await _listener(ep, cache_ep=True) + + async def _send_op(op, port, ep=None): op_msg = generate_op_message(op, port) msg2send = np.arange(10) msg2recv = np.empty_like(msg2send) - ep = await ucp.create_endpoint(ucp.get_address(), port) + if ep is None: + ep = await create_endpoint_retry(monitor_port, port, "Monitor", "Monitor") msgs = [ep.send(op_msg), ep.send(msg2send), ep.recv(msg2recv)] await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) # Start monitor's listener - listener = ucp.create_listener(_listener, port=monitor_port) + listener = ucp.create_listener(_listener_cb, port=monitor_port) # Wait until all workers signal they are listening while len(listening_worker_ports) != len(worker_ports): await asyncio.sleep(0.1) - # Send cluster ready message to all workers + # Create persistent endpoints to all workers + worker_eps = {} + if PersistentEndpoints: + for remote_port in worker_ports: + worker_eps[remote_port] = await create_endpoint_retry( + monitor_port, remote_port, "Monitor", "Worker" + ) + + # Send shutdown message to all workers ready_signals = [] for port in listening_worker_ports: - ready_signals.append(_send_op(OP_CLUSTER_READY, port)) + if PersistentEndpoints: + ready_signals.append(_send_op(OP_CLUSTER_READY, port, worker_eps[port])) + else: + ready_signals.append(_send_op(OP_CLUSTER_READY, port)) await asyncio.gather(*ready_signals, loop=asyncio.get_event_loop()) + # When using persistent endpoints, we need to wait on previously + # created endpoints for completion signal + if PersistentEndpoints: + listener_tasks = [] + for listener_ep in listener_eps: + listener_tasks.append(_listener(listener_ep)) + + await asyncio.gather(*listener_tasks, loop=asyncio.get_event_loop()) + # Wait until all workers signal completion while len(completed_worker_ports) != len(worker_ports): await asyncio.sleep(0.1) @@ -162,7 +263,10 @@ async def _send_op(op, port): # Send shutdown message to all workers close = [] for port in completed_worker_ports: - close.append(_send_op(OP_SHUTDOWN, port)) + if PersistentEndpoints: + close.append(_send_op(OP_SHUTDOWN, port, ep=worker_eps[port])) + else: + close.append(_send_op(OP_SHUTDOWN, port)) await asyncio.gather(*close, loop=asyncio.get_event_loop()) listener.close() From 65f1870c8ba5cf32402a816d371b8260f78f6570 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Wed, 10 Mar 2021 09:42:31 -0800 Subject: [PATCH 03/42] Improve closing of endpoints in multiproc all_to_all test --- tests/test_multiple_processes_all_to_all.py | 44 ++++++++++++--------- 1 file changed, 25 insertions(+), 19 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index e5618fb13..f352f9383 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -1,5 +1,6 @@ import asyncio import multiprocessing +import os import random import sys @@ -17,7 +18,7 @@ OP_CLUSTER_READY = 3 OP_SHUTDOWN = 4 -PersistentEndpoints = False +PersistentEndpoints = True def generate_op_message(op, port): @@ -49,6 +50,7 @@ async def create_endpoint_retry(my_port, remote_port, my_task, remote_task): def worker(my_port, monitor_port, all_ports): ucp.init() + eps = [] listener_eps = [] global listener_monitor_ep @@ -62,11 +64,9 @@ def _register_cluster_started(): global cluster_started cluster_started = True - async def _close_endpoints(): - for ep in listener_eps: - await ep.close() - async def _listener(ep, cache_ep=False): + global listener_monitor_ep + op_msg = generate_op_message(OP_NONE, 0) msg2send = np.arange(10) msg2recv = np.empty_like(msg2send) @@ -76,19 +76,20 @@ async def _listener(ep, cache_ep=False): op = parse_op_message(op_msg)["op"] + if op == OP_SHUTDOWN: + while not listener_monitor_ep.closed(): + await asyncio.sleep(0.1) + listener.close() + return + if op == OP_CLUSTER_READY: + _register_cluster_started() + if cache_ep and PersistentEndpoints: if op == OP_NONE: listener_eps.append(ep) else: - global listener_monitor_ep listener_monitor_ep = ep - if op == OP_SHUTDOWN: - await ep.close() - listener.close() - if op == OP_CLUSTER_READY: - _register_cluster_started() - async def _listener_cb(ep): await _listener(ep, cache_ep=True) @@ -107,8 +108,9 @@ async def _signal_monitor(monitor_port, my_port, op, ep=None): ack_msg = bytearray(2) if ep is None: - #ep = await ucp.create_endpoint(ucp.get_address(), monitor_port) - ep = await create_endpoint_retry(my_port, monitor_port, "Worker", "Monitor") + ep = await create_endpoint_retry( + my_port, monitor_port, "Worker", "Monitor" + ) msgs = [ep.send(op_msg), ep.recv(ack_msg)] await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) @@ -130,7 +132,6 @@ async def _signal_monitor(monitor_port, my_port, op, ep=None): while not cluster_started: await asyncio.sleep(0.1) - eps = [] if PersistentEndpoints: client_tasks = [] # Create endpoints to all other workers @@ -181,7 +182,7 @@ async def _signal_monitor(monitor_port, my_port, op, ep=None): def monitor(monitor_port, worker_ports): ucp.init() - listener_eps = [] + listener_eps = {} listening_worker_ports = [] completed_worker_ports = [] @@ -194,7 +195,7 @@ def _register(op, port): async def _listener(ep, cache_ep=True): if cache_ep and PersistentEndpoints: - listener_eps.append(ep) + listener_eps[ep.uid] = ep op_msg = generate_op_message(OP_NONE, 0) ack_msg = bytearray(int(888).to_bytes(2, sys.byteorder)) @@ -219,10 +220,15 @@ async def _send_op(op, port, ep=None): msg2recv = np.empty_like(msg2send) if ep is None: - ep = await create_endpoint_retry(monitor_port, port, "Monitor", "Monitor") + ep = await create_endpoint_retry( + monitor_port, port, "Monitor", "Monitor" + ) msgs = [ep.send(op_msg), ep.send(msg2send), ep.recv(msg2recv)] await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) + if op == OP_SHUTDOWN: + await ep.close() + # Start monitor's listener listener = ucp.create_listener(_listener_cb, port=monitor_port) @@ -251,7 +257,7 @@ async def _send_op(op, port, ep=None): # created endpoints for completion signal if PersistentEndpoints: listener_tasks = [] - for listener_ep in listener_eps: + for listener_ep in listener_eps.values(): listener_tasks.append(_listener(listener_ep)) await asyncio.gather(*listener_tasks, loop=asyncio.get_event_loop()) From ba66c92498f5866de03d65eda33db51cbc3def75 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Wed, 10 Mar 2021 10:11:54 -0800 Subject: [PATCH 04/42] Do more transfers between worker pairs --- tests/test_multiple_processes_all_to_all.py | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index f352f9383..a90684337 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -148,6 +148,18 @@ async def _signal_monitor(monitor_port, my_port, op, ep=None): # Wait until listener_eps have all been cached while len(listener_eps) != len(all_ports) - 1: await asyncio.sleep(0.1) + + # Exchange messages with other workers + for i in range(3): + client_tasks = [] + listener_tasks = [] + for ep in eps: + client_tasks.append(_client(remote_port, ep)) + for listener_ep in listener_eps: + listener_tasks.append(_listener(listener_ep)) + + all_tasks = client_tasks + listener_tasks + await asyncio.gather(*all_tasks, loop=asyncio.get_event_loop()) else: # Create endpoints to all other workers client_tasks = [] From b712ce9d99b6a2c7e726be47b5e62d797c5d4ba9 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Wed, 10 Mar 2021 14:54:50 -0800 Subject: [PATCH 05/42] Support multiple endpoints per worker in multiproc all_to_all test --- tests/test_multiple_processes_all_to_all.py | 34 ++++++++++++--------- 1 file changed, 19 insertions(+), 15 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index a90684337..5194d5d3e 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -47,7 +47,7 @@ async def create_endpoint_retry(my_port, remote_port, my_task, remote_task): await asyncio.sleep(0.1) -def worker(my_port, monitor_port, all_ports): +def worker(my_port, monitor_port, all_ports, endpoints_per_worker): ucp.init() eps = [] @@ -133,20 +133,21 @@ async def _signal_monitor(monitor_port, my_port, op, ep=None): await asyncio.sleep(0.1) if PersistentEndpoints: - client_tasks = [] - # Create endpoints to all other workers - for remote_port in all_ports: - if remote_port == my_port: - continue - ep = await create_endpoint_retry( - my_port, remote_port, "Worker", "Worker" - ) - eps.append(ep) - client_tasks.append(_client(remote_port, ep)) - await asyncio.gather(*client_tasks, loop=asyncio.get_event_loop()) + for i in range(endpoints_per_worker): + client_tasks = [] + # Create endpoints to all other workers + for remote_port in all_ports: + if remote_port == my_port: + continue + ep = await create_endpoint_retry( + my_port, remote_port, "Worker", "Worker" + ) + eps.append(ep) + client_tasks.append(_client(remote_port, ep)) + await asyncio.gather(*client_tasks, loop=asyncio.get_event_loop()) # Wait until listener_eps have all been cached - while len(listener_eps) != len(all_ports) - 1: + while len(listener_eps) != endpoints_per_worker * (len(all_ports) - 1): await asyncio.sleep(0.1) # Exchange messages with other workers @@ -293,7 +294,8 @@ async def _send_op(op, port, ep=None): @pytest.mark.parametrize("num_workers", [1, 2, 4, 8]) -def test_send_recv_cu(num_workers): +@pytest.mark.parametrize("endpoints_per_worker", [20, 80, 320, 640]) +def test_send_recv_cu(num_workers, endpoints_per_worker): # One additional port for monitor num_ports = num_workers + 1 @@ -318,7 +320,9 @@ def test_send_recv_cu(num_workers): worker_processes = [] for port in worker_ports: worker_process = ctx.Process( - name="worker", target=worker, args=[port, monitor_port, worker_ports] + name="worker", + target=worker, + args=[port, monitor_port, worker_ports, endpoints_per_worker], ) worker_process.start() worker_processes.append(worker_process) From aa067315737b4df91d89e2571c2ea15262f39154 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Wed, 9 Jun 2021 15:54:17 -0700 Subject: [PATCH 06/42] Remove monitor process in favor of multiprocessing shared memory --- tests/test_multiple_processes_all_to_all.py | 270 +++----------------- 1 file changed, 42 insertions(+), 228 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index 5194d5d3e..922197229 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -1,38 +1,14 @@ import asyncio import multiprocessing -import os -import random -import sys import numpy as np import pytest import ucp -OP_BYTES = 1 -PORT_BYTES = 2 - -OP_NONE = 0 -OP_WORKER_LISTENING = 1 -OP_WORKER_COMPLETED = 2 -OP_CLUSTER_READY = 3 -OP_SHUTDOWN = 4 - PersistentEndpoints = True -def generate_op_message(op, port): - op_bytes = op.to_bytes(OP_BYTES, sys.byteorder) - port_bytes = port.to_bytes(PORT_BYTES, sys.byteorder) - return bytearray(b"".join([op_bytes, port_bytes])) - - -def parse_op_message(msg): - op = int.from_bytes(msg[0:OP_BYTES], sys.byteorder) - port = int.from_bytes(msg[OP_BYTES : OP_BYTES + PORT_BYTES], sys.byteorder) - return {"op": op, "port": port} - - async def create_endpoint_retry(my_port, remote_port, my_task, remote_task): while True: try: @@ -47,107 +23,66 @@ async def create_endpoint_retry(my_port, remote_port, my_task, remote_task): await asyncio.sleep(0.1) -def worker(my_port, monitor_port, all_ports, endpoints_per_worker): +def worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): ucp.init() eps = [] - listener_eps = [] - - global listener_monitor_ep - listener_monitor_ep = None + listener_eps = set() global cluster_started cluster_started = False - async def _worker(my_port, all_ports): + async def _worker(): def _register_cluster_started(): global cluster_started cluster_started = True async def _listener(ep, cache_ep=False): - global listener_monitor_ep - - op_msg = generate_op_message(OP_NONE, 0) msg2send = np.arange(10) msg2recv = np.empty_like(msg2send) - msgs = [ep.recv(op_msg), ep.send(msg2send), ep.recv(msg2recv)] + msgs = [ep.send(msg2send), ep.recv(msg2recv)] await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) - op = parse_op_message(op_msg)["op"] - - if op == OP_SHUTDOWN: - while not listener_monitor_ep.closed(): - await asyncio.sleep(0.1) - listener.close() - return - if op == OP_CLUSTER_READY: - _register_cluster_started() - - if cache_ep and PersistentEndpoints: - if op == OP_NONE: - listener_eps.append(ep) - else: - listener_monitor_ep = ep - async def _listener_cb(ep): + if PersistentEndpoints: + listener_eps.add(ep) await _listener(ep, cache_ep=True) - async def _client(port, ep=None): - op_msg = generate_op_message(OP_NONE, 0) + async def _client(my_port, remote_port, ep=None): msg2send = np.arange(10) msg2recv = np.empty_like(msg2send) if ep is None: ep = await create_endpoint_retry(my_port, port, "Worker", "Worker") - msgs = [ep.send(op_msg), ep.recv(msg2recv), ep.send(msg2send)] - await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) - - async def _signal_monitor(monitor_port, my_port, op, ep=None): - op_msg = generate_op_message(op, my_port) - ack_msg = bytearray(2) - - if ep is None: - ep = await create_endpoint_retry( - my_port, monitor_port, "Worker", "Monitor" - ) - msgs = [ep.send(op_msg), ep.recv(ack_msg)] + msgs = [ep.recv(msg2recv), ep.send(msg2send)] await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) # Start listener - listener = ucp.create_listener(_listener_cb, port=my_port) + listener = ucp.create_listener(_listener_cb) + with lock: + signal[0] += 1 + ports[worker_num] = listener.port - # Signal monitor that worker is listening - monitor_ep = None - if PersistentEndpoints: - monitor_ep = await create_endpoint_retry( - my_port, monitor_port, "Worker", "Monitor" - ) - await _signal_monitor( - monitor_port, my_port, op=OP_WORKER_LISTENING, ep=monitor_ep - ) - else: - await _signal_monitor(monitor_port, my_port, op=OP_WORKER_LISTENING) - - while not cluster_started: - await asyncio.sleep(0.1) + while signal[0] != num_workers: + pass if PersistentEndpoints: for i in range(endpoints_per_worker): client_tasks = [] # Create endpoints to all other workers - for remote_port in all_ports: - if remote_port == my_port: + for remote_port in list(ports): + if remote_port == listener.port: continue ep = await create_endpoint_retry( - my_port, remote_port, "Worker", "Worker" + listener.port, remote_port, "Worker", "Worker" ) eps.append(ep) - client_tasks.append(_client(remote_port, ep)) + client_tasks.append(_client(listener.port, remote_port, ep)) await asyncio.gather(*client_tasks, loop=asyncio.get_event_loop()) # Wait until listener_eps have all been cached - while len(listener_eps) != endpoints_per_worker * (len(all_ports) - 1): + while len(listener_eps) != endpoints_per_worker * (num_workers - 1): await asyncio.sleep(0.1) # Exchange messages with other workers @@ -155,32 +90,30 @@ async def _signal_monitor(monitor_port, my_port, op, ep=None): client_tasks = [] listener_tasks = [] for ep in eps: - client_tasks.append(_client(remote_port, ep)) + client_tasks.append(_client(listener.port, remote_port, ep)) for listener_ep in listener_eps: listener_tasks.append(_listener(listener_ep)) all_tasks = client_tasks + listener_tasks await asyncio.gather(*all_tasks, loop=asyncio.get_event_loop()) else: - # Create endpoints to all other workers - client_tasks = [] - for port in all_ports: - if port == my_port: - continue - client_tasks.append(_client(port)) - await asyncio.gather(*client_tasks, loop=asyncio.get_event_loop()) - - # Signal monitor that worker is completed - if PersistentEndpoints: - await _signal_monitor( - monitor_port, my_port, op=OP_WORKER_COMPLETED, ep=monitor_ep - ) - else: - await _signal_monitor(monitor_port, my_port, op=OP_WORKER_COMPLETED) + for i in range(3): + # Create endpoints to all other workers + client_tasks = [] + for port in list(ports): + if port == listener.port: + continue + client_tasks.append(_client(listener.port, port)) + await asyncio.gather(*client_tasks, loop=asyncio.get_event_loop()) - # Wait for closing signal - if PersistentEndpoints: - await _listener(listener_monitor_ep) + with lock: + signal[1] += 1 + ports[worker_num] = listener.port + + while signal[1] != num_workers: + pass + + listener.close() # Wait for a shutdown signal from monitor try: @@ -189,140 +122,24 @@ async def _signal_monitor(monitor_port, my_port, op, ep=None): except ucp.UCXCloseError: pass - asyncio.get_event_loop().run_until_complete(_worker(my_port, all_ports)) - - -def monitor(monitor_port, worker_ports): - ucp.init() - - listener_eps = {} - listening_worker_ports = [] - completed_worker_ports = [] - - async def _monitor(monitor_port, worker_ports): - def _register(op, port): - if op == OP_WORKER_LISTENING: - listening_worker_ports.append(port) - elif op == OP_WORKER_COMPLETED: - completed_worker_ports.append(port) - - async def _listener(ep, cache_ep=True): - if cache_ep and PersistentEndpoints: - listener_eps[ep.uid] = ep - - op_msg = generate_op_message(OP_NONE, 0) - ack_msg = bytearray(int(888).to_bytes(2, sys.byteorder)) - - # Sending an ack_msg prevents the other ep from closing too - # early, ultimately leading this process to hang. - msgs = [ep.recv(op_msg), ep.send(ack_msg)] - await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) - - op_msg = parse_op_message(op_msg) - worker_op = op_msg["op"] - worker_port = op_msg["port"] - - _register(worker_op, worker_port) - - async def _listener_cb(ep): - await _listener(ep, cache_ep=True) - - async def _send_op(op, port, ep=None): - op_msg = generate_op_message(op, port) - msg2send = np.arange(10) - msg2recv = np.empty_like(msg2send) - - if ep is None: - ep = await create_endpoint_retry( - monitor_port, port, "Monitor", "Monitor" - ) - msgs = [ep.send(op_msg), ep.send(msg2send), ep.recv(msg2recv)] - await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) - - if op == OP_SHUTDOWN: - await ep.close() - - # Start monitor's listener - listener = ucp.create_listener(_listener_cb, port=monitor_port) - - # Wait until all workers signal they are listening - while len(listening_worker_ports) != len(worker_ports): - await asyncio.sleep(0.1) - - # Create persistent endpoints to all workers - worker_eps = {} - if PersistentEndpoints: - for remote_port in worker_ports: - worker_eps[remote_port] = await create_endpoint_retry( - monitor_port, remote_port, "Monitor", "Worker" - ) - - # Send shutdown message to all workers - ready_signals = [] - for port in listening_worker_ports: - if PersistentEndpoints: - ready_signals.append(_send_op(OP_CLUSTER_READY, port, worker_eps[port])) - else: - ready_signals.append(_send_op(OP_CLUSTER_READY, port)) - await asyncio.gather(*ready_signals, loop=asyncio.get_event_loop()) - - # When using persistent endpoints, we need to wait on previously - # created endpoints for completion signal - if PersistentEndpoints: - listener_tasks = [] - for listener_ep in listener_eps.values(): - listener_tasks.append(_listener(listener_ep)) - - await asyncio.gather(*listener_tasks, loop=asyncio.get_event_loop()) - - # Wait until all workers signal completion - while len(completed_worker_ports) != len(worker_ports): - await asyncio.sleep(0.1) - - # Send shutdown message to all workers - close = [] - for port in completed_worker_ports: - if PersistentEndpoints: - close.append(_send_op(OP_SHUTDOWN, port, ep=worker_eps[port])) - else: - close.append(_send_op(OP_SHUTDOWN, port)) - await asyncio.gather(*close, loop=asyncio.get_event_loop()) - - listener.close() - - asyncio.get_event_loop().run_until_complete(_monitor(monitor_port, worker_ports)) + asyncio.get_event_loop().run_until_complete(_worker()) @pytest.mark.parametrize("num_workers", [1, 2, 4, 8]) @pytest.mark.parametrize("endpoints_per_worker", [20, 80, 320, 640]) def test_send_recv_cu(num_workers, endpoints_per_worker): - # One additional port for monitor - num_ports = num_workers + 1 - - ports = set() - while len(ports) != num_ports: - missing_ports = num_ports - len(ports) - ports = ports.union( - [random.randint(13000, 23000) for n in range(missing_ports)] - ) - ports = list(ports) - - monitor_port = ports[0] - worker_ports = ports[1:] - ctx = multiprocessing.get_context("spawn") - monitor_process = ctx.Process( - name="monitor", target=monitor, args=[monitor_port, worker_ports] - ) - monitor_process.start() + signal = ctx.Array("i", [0, 0]) + ports = ctx.Array("i", range(num_workers)) + lock = ctx.Lock() worker_processes = [] - for port in worker_ports: + for worker_num in range(num_workers): worker_process = ctx.Process( name="worker", target=worker, - args=[port, monitor_port, worker_ports, endpoints_per_worker], + args=[signal, ports, lock, worker_num, num_workers, endpoints_per_worker], ) worker_process.start() worker_processes.append(worker_process) @@ -330,7 +147,4 @@ def test_send_recv_cu(num_workers, endpoints_per_worker): for worker_process in worker_processes: worker_process.join() - monitor_process.join() - assert worker_process.exitcode == 0 - assert monitor_process.exitcode == 0 From 68c07a6d173c850c83e1505e0808400d5224ff1d Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Thu, 1 Jul 2021 11:39:43 -0700 Subject: [PATCH 07/42] Mark some multiple processes all-to-all tests as slow --- tests/test_multiple_processes_all_to_all.py | 17 ++++++++++++++--- 1 file changed, 14 insertions(+), 3 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index 922197229..f2eb45334 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -125,9 +125,7 @@ async def _client(my_port, remote_port, ep=None): asyncio.get_event_loop().run_until_complete(_worker()) -@pytest.mark.parametrize("num_workers", [1, 2, 4, 8]) -@pytest.mark.parametrize("endpoints_per_worker", [20, 80, 320, 640]) -def test_send_recv_cu(num_workers, endpoints_per_worker): +def _test_send_recv_cu(num_workers, endpoints_per_worker): ctx = multiprocessing.get_context("spawn") signal = ctx.Array("i", [0, 0]) @@ -148,3 +146,16 @@ def test_send_recv_cu(num_workers, endpoints_per_worker): worker_process.join() assert worker_process.exitcode == 0 + + +@pytest.mark.parametrize("num_workers", [1, 2, 4, 8]) +@pytest.mark.parametrize("endpoints_per_worker", [20, 80]) +def test_send_recv_cu(num_workers, endpoints_per_worker): + _test_send_recv_cu(num_workers, endpoints_per_worker) + + +@pytest.mark.slow +@pytest.mark.parametrize("num_workers", [1, 2, 4, 8]) +@pytest.mark.parametrize("endpoints_per_worker", [320, 640]) +def test_send_recv_cu_slow(num_workers, endpoints_per_worker): + _test_send_recv_cu(num_workers, endpoints_per_worker) From c91ec3cc4ab5807979764c12fb0b3ab6dff05d4c Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Thu, 1 Jul 2021 14:00:05 -0700 Subject: [PATCH 08/42] Store remote EP port in dictionary --- tests/test_multiple_processes_all_to_all.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index f2eb45334..b69b34e7c 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -26,7 +26,7 @@ async def create_endpoint_retry(my_port, remote_port, my_task, remote_task): def worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): ucp.init() - eps = [] + eps = dict() listener_eps = set() global cluster_started @@ -77,7 +77,7 @@ async def _client(my_port, remote_port, ep=None): ep = await create_endpoint_retry( listener.port, remote_port, "Worker", "Worker" ) - eps.append(ep) + eps[(remote_port, i)] = ep client_tasks.append(_client(listener.port, remote_port, ep)) await asyncio.gather(*client_tasks, loop=asyncio.get_event_loop()) @@ -89,7 +89,7 @@ async def _client(my_port, remote_port, ep=None): for i in range(3): client_tasks = [] listener_tasks = [] - for ep in eps: + for (remote_port, _), ep in eps.items(): client_tasks.append(_client(listener.port, remote_port, ep)) for listener_ep in listener_eps: listener_tasks.append(_listener(listener_ep)) From 33b630df3d04964709ab28e8d5ca390e9e026c70 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Thu, 1 Jul 2021 14:00:42 -0700 Subject: [PATCH 09/42] Mark more tests as slow --- tests/test_multiple_processes_all_to_all.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index b69b34e7c..2bf5e1ef6 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -149,13 +149,13 @@ def _test_send_recv_cu(num_workers, endpoints_per_worker): @pytest.mark.parametrize("num_workers", [1, 2, 4, 8]) -@pytest.mark.parametrize("endpoints_per_worker", [20, 80]) +@pytest.mark.parametrize("endpoints_per_worker", [20, 40]) def test_send_recv_cu(num_workers, endpoints_per_worker): _test_send_recv_cu(num_workers, endpoints_per_worker) @pytest.mark.slow @pytest.mark.parametrize("num_workers", [1, 2, 4, 8]) -@pytest.mark.parametrize("endpoints_per_worker", [320, 640]) +@pytest.mark.parametrize("endpoints_per_worker", [80, 320, 640]) def test_send_recv_cu_slow(num_workers, endpoints_per_worker): _test_send_recv_cu(num_workers, endpoints_per_worker) From 4f41ebeee5f9457f26f42e5c07e0eece9e7ccc05 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Thu, 1 Jul 2021 14:45:14 -0700 Subject: [PATCH 10/42] Benchmark for multiprocess all-to-all --- tests/test_multiple_processes_all_to_all.py | 75 +++++++++++++++------ 1 file changed, 56 insertions(+), 19 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index 2bf5e1ef6..5fd36bb4d 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -1,12 +1,19 @@ import asyncio import multiprocessing +from time import monotonic import numpy as np import pytest +from dask.utils import format_bytes +from distributed.utils import nbytes + import ucp PersistentEndpoints = True +GatherAsync = False +Iterations = 3 +Size = 2 ** 25 async def create_endpoint_retry(my_port, remote_port, my_task, remote_task): @@ -28,6 +35,7 @@ def worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): eps = dict() listener_eps = set() + bytes_bandwidth = dict() global cluster_started cluster_started = False @@ -37,26 +45,43 @@ def _register_cluster_started(): global cluster_started cluster_started = True - async def _listener(ep, cache_ep=False): - msg2send = np.arange(10) + async def _listener(ep): + msg2send = np.arange(Size) msg2recv = np.empty_like(msg2send) - msgs = [ep.send(msg2send), ep.recv(msg2recv)] - await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) + if GatherAsync: + msgs = [ep.send(msg2send), ep.recv(msg2recv)] * Iterations + await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) + else: + for i in range(Iterations): + msgs = [ep.send(msg2send), ep.recv(msg2recv)] + await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) async def _listener_cb(ep): if PersistentEndpoints: listener_eps.add(ep) - await _listener(ep, cache_ep=True) + await _listener(ep) - async def _client(my_port, remote_port, ep=None): - msg2send = np.arange(10) + async def _client(my_port, remote_port, ep=None, cache_only=False): + msg2send = np.arange(Size) msg2recv = np.empty_like(msg2send) + send_recv_bytes = (nbytes(msg2send) + nbytes(msg2recv)) * Iterations if ep is None: ep = await create_endpoint_retry(my_port, port, "Worker", "Worker") - msgs = [ep.recv(msg2recv), ep.send(msg2send)] - await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) + + t = monotonic() + if GatherAsync: + msgs = [ep.recv(msg2recv), ep.send(msg2send)] * Iterations + await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) + else: + for i in range(Iterations): + msgs = [ep.recv(msg2recv), ep.send(msg2send)] * Iterations + await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) + if cache_only is False: + bytes_bandwidth[remote_port].append( + (send_recv_bytes, send_recv_bytes / (monotonic() - t)) + ) # Start listener listener = ucp.create_listener(_listener_cb) @@ -77,8 +102,11 @@ async def _client(my_port, remote_port, ep=None): ep = await create_endpoint_retry( listener.port, remote_port, "Worker", "Worker" ) + bytes_bandwidth[remote_port] = [] eps[(remote_port, i)] = ep - client_tasks.append(_client(listener.port, remote_port, ep)) + client_tasks.append( + _client(listener.port, remote_port, ep, cache_only=True) + ) await asyncio.gather(*client_tasks, loop=asyncio.get_event_loop()) # Wait until listener_eps have all been cached @@ -113,6 +141,22 @@ async def _client(my_port, remote_port, ep=None): while signal[1] != num_workers: pass + for remote_port, bb in bytes_bandwidth.items(): + 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( + "[%d, %d] Transferred bytes: %s, average bandwidth: %s/s, " + "median bandwidth: %s/s" + % ( + listener.port, + remote_port, + format_bytes(total_bytes), + format_bytes(avg_bandwidth), + format_bytes(median_bandwidth), + ) + ) + listener.close() # Wait for a shutdown signal from monitor @@ -148,14 +192,7 @@ def _test_send_recv_cu(num_workers, endpoints_per_worker): assert worker_process.exitcode == 0 -@pytest.mark.parametrize("num_workers", [1, 2, 4, 8]) -@pytest.mark.parametrize("endpoints_per_worker", [20, 40]) +@pytest.mark.parametrize("num_workers", [2, 4, 8]) +@pytest.mark.parametrize("endpoints_per_worker", [1]) def test_send_recv_cu(num_workers, endpoints_per_worker): _test_send_recv_cu(num_workers, endpoints_per_worker) - - -@pytest.mark.slow -@pytest.mark.parametrize("num_workers", [1, 2, 4, 8]) -@pytest.mark.parametrize("endpoints_per_worker", [80, 320, 640]) -def test_send_recv_cu_slow(num_workers, endpoints_per_worker): - _test_send_recv_cu(num_workers, endpoints_per_worker) From b6508ed4cc7bc437779a93a17c05cc6cab02b133 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Mon, 5 Jul 2021 13:50:48 -0700 Subject: [PATCH 11/42] Add Tornado-based all-to-all benchmark --- tests/test_multiple_processes_all_to_all.py | 168 +++++++++++++++++++- tests/utils.py | 136 ++++++++++++++++ 2 files changed, 301 insertions(+), 3 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index 5fd36bb4d..e750bdb28 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -4,8 +4,13 @@ import numpy as np import pytest +from tornado import gen +from tornado.ioloop import IOLoop +from utils import TornadoTCPConnection, TornadoTCPServer from dask.utils import format_bytes +from distributed.comm.utils import to_frames +from distributed.protocol import to_serialize from distributed.utils import nbytes import ucp @@ -13,7 +18,7 @@ PersistentEndpoints = True GatherAsync = False Iterations = 3 -Size = 2 ** 25 +Size = 2 ** 15 async def create_endpoint_retry(my_port, remote_port, my_task, remote_task): @@ -76,7 +81,7 @@ async def _client(my_port, remote_port, ep=None, cache_only=False): await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) else: for i in range(Iterations): - msgs = [ep.recv(msg2recv), ep.send(msg2send)] * Iterations + msgs = [ep.recv(msg2recv), ep.send(msg2send)] await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) if cache_only is False: bytes_bandwidth[remote_port].append( @@ -169,6 +174,161 @@ async def _client(my_port, remote_port, ep=None, cache_only=False): asyncio.get_event_loop().run_until_complete(_worker()) +def tornado_worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): + conns = dict() + bytes_bandwidth = dict() + + global cluster_started + cluster_started = False + + async def _worker(): + def _register_cluster_started(): + global cluster_started + cluster_started = True + + async def _get_message_size_and_frames(): + message = np.arange(Size) + + msg = {"data": to_serialize(message)} + frames = await to_frames(msg, serializers=("cuda", "dask", "pickle")) + + return (nbytes(message), frames) + + async def _listener(conn): + _, frames = await _get_message_size_and_frames() + + if GatherAsync: + msgs = [ + op + for i in range(Iterations) + for op in [conn.recv(), conn.send(frames)] + ] + await gen.multi(msgs) + else: + for i in range(Iterations): + msgs = [conn.send(frames), conn.recv()] + await gen.multi(msgs) + + # This seems to be faster! + # await conn.send(frames) + # await conn.recv() + + async def _client(my_port, remote_port, conn=None): + message_size, frames = await _get_message_size_and_frames() + send_recv_bytes = (message_size * 2) * Iterations + + # if ep is None: + # ep = await create_endpoint_retry(my_port, port, "Worker", "Worker") + + t = monotonic() + if GatherAsync: + msgs = [ + op + for i in range(Iterations) + for op in [conn.recv(), conn.send(frames)] + ] + await gen.multi(msgs) + else: + for i in range(Iterations): + msgs = [conn.recv(), conn.send(frames)] + await gen.multi(msgs) + + # This seems to be faster! + # await conn.recv() + # await conn.send(frames) + + bytes_bandwidth[remote_port].append( + (send_recv_bytes, send_recv_bytes / (monotonic() - t)) + ) + + host = ucp.get_address(ifname="enp1s0f0") + + # Start listener + listener = await TornadoTCPServer.start_server(host, None) + with lock: + signal[0] += 1 + ports[worker_num] = listener.port + + while signal[0] != num_workers: + await gen.sleep(0) + + print(list(ports)) + + if PersistentEndpoints: + for i in range(endpoints_per_worker): + client_tasks = [] + # Create endpoints to all other workers + for remote_port in list(ports): + if remote_port == listener.port: + continue + conn = await TornadoTCPConnection.connect(host, remote_port) + conns[(remote_port, i)] = conn + bytes_bandwidth[remote_port] = [] + + # Wait until all clients connected to listener + while len(listener.get_connections()) != endpoints_per_worker * ( + num_workers - 1 + ): + await gen.sleep(0) + + # Exchange messages with other workers + for i in range(3): + client_tasks = [] + listener_tasks = [] + for (remote_port, _), conn in conns.items(): + client_tasks.append(_client(listener.port, remote_port, conn)) + for conn in listener.get_connections(): + listener_tasks.append(_listener(conn)) + + all_tasks = client_tasks + listener_tasks + await gen.multi(all_tasks) + else: + for i in range(3): + # Create endpoints to all other workers + client_tasks = [] + for port in list(ports): + if port == listener.port: + continue + client_tasks.append(_client(listener.port, port)) + await gen.multi(client_tasks) + + with lock: + signal[1] += 1 + ports[worker_num] = listener.port + + while signal[1] != num_workers: + pass + + for remote_port, bb in bytes_bandwidth.items(): + 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( + "[%d, %d] Transferred bytes: %s, average bandwidth: %s/s, " + "median bandwidth: %s/s" + % ( + listener.port, + remote_port, + format_bytes(total_bytes), + format_bytes(avg_bandwidth), + format_bytes(median_bandwidth), + ) + ) + + # listener.server.close() + for conn in listener.get_connections(): + conn.stream.close() + + # Wait for a shutdown signal from monitor + try: + while not all(c.stream.closed() for c in listener.get_connections()): + await gen.sleep(0) + except ucp.UCXCloseError: + pass + + IOLoop.current().run_sync(_worker) + + def _test_send_recv_cu(num_workers, endpoints_per_worker): ctx = multiprocessing.get_context("spawn") @@ -180,7 +340,9 @@ def _test_send_recv_cu(num_workers, endpoints_per_worker): for worker_num in range(num_workers): worker_process = ctx.Process( name="worker", - target=worker, + # target=worker, + # target=asyncio_worker, + target=tornado_worker, args=[signal, ports, lock, worker_num, num_workers, endpoints_per_worker], ) worker_process.start() diff --git a/tests/utils.py b/tests/utils.py index ef8ecccd7..23aeea836 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -1,11 +1,16 @@ import io import logging import os +import struct from contextlib import contextmanager 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 +from distributed.protocol.utils import pack_frames_prelude, unpack_frames from distributed.utils import nbytes import rmm @@ -138,3 +143,134 @@ async def am_recv(ep): msg = await from_frames(frames) return frames, msg + + +class TornadoTCPConnection: + def __init__(self, stream, client=None): + self._client = client + self.stream = stream + + @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, frames): + stream = self.stream + if stream is None: + raise StreamClosedError() + + 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 msg + + +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: + + 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 From ac880607e70ec46a1e19f11d88f44b400d70932f Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Tue, 6 Jul 2021 03:12:10 -0700 Subject: [PATCH 12/42] Change all-to-all worker function into class --- tests/test_multiple_processes_all_to_all.py | 233 +++++++++++--------- 1 file changed, 130 insertions(+), 103 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index e750bdb28..33dcd20fe 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -21,132 +21,167 @@ Size = 2 ** 15 -async def create_endpoint_retry(my_port, remote_port, my_task, remote_task): - while True: - try: - ep = await ucp.create_endpoint(ucp.get_address(), remote_port) - return ep - except ucp.exceptions.UCXCanceled as e: - print( - "%s[%d]->%s[%d] Failed: %s" - % (my_task, my_port, remote_task, remote_port, e), - flush=True, - ) - await asyncio.sleep(0.1) - +class BaseWorker: + def __init__( + self, signal, ports, lock, worker_num, num_workers, endpoints_per_worker + ): + self.signal = signal + self.ports = ports + self.lock = lock + self.worker_num = worker_num + self.num_workers = num_workers + self.endpoints_per_worker = endpoints_per_worker + + 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, msg2recv, send_first=True): + if GatherAsync: + msgs = [ep.send(msg2send), ep.recv(msg2recv)] * Iterations + await self._gather(msgs) + else: + for i in range(Iterations): + msgs = [ep.send(msg2send), ep.recv(msg2recv)] + await self._gather(msgs) -def worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): - ucp.init() + async def _listener(self, ep): + msg2send = np.arange(Size) + msg2recv = np.empty_like(msg2send) - eps = dict() - listener_eps = set() - bytes_bandwidth = dict() + await self._transfer(ep, msg2send, msg2recv) - global cluster_started - cluster_started = False + async def _client(self, my_port, remote_port, ep=None, cache_only=False): + msg2send = np.arange(Size) + msg2recv = np.empty_like(msg2send) + send_recv_bytes = (nbytes(msg2send) + nbytes(msg2recv)) * Iterations - async def _worker(): - def _register_cluster_started(): - global cluster_started - cluster_started = True + if ep is None: + ep = await self._create_endpoint(remote_port) - async def _listener(ep): - msg2send = np.arange(Size) - msg2recv = np.empty_like(msg2send) + t = monotonic() + await self._transfer(ep, msg2send, msg2recv, send_first=False) + total_time = monotonic() - t - if GatherAsync: - msgs = [ep.send(msg2send), ep.recv(msg2recv)] * Iterations - await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) - else: - for i in range(Iterations): - msgs = [ep.send(msg2send), ep.recv(msg2recv)] - await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) + if cache_only is False: + self.bytes_bandwidth[remote_port].append( + (send_recv_bytes, send_recv_bytes / total_time) + ) - async def _listener_cb(ep): - if PersistentEndpoints: - listener_eps.add(ep) - await _listener(ep) + def _init(self): + ucp.init() - async def _client(my_port, remote_port, ep=None, cache_only=False): - msg2send = np.arange(Size) - msg2recv = np.empty_like(msg2send) - send_recv_bytes = (nbytes(msg2send) + nbytes(msg2recv)) * Iterations + async def _listener_cb(self, ep): + if PersistentEndpoints: + self.connections.add(ep) + await self._listener(ep) + + def _create_listener(self): + return ucp.create_listener(self._listener_cb) + + async def _create_endpoint(self, remote_port): + my_port = self.listener_port + my_task = "Worker" + remote_task = "Worker" + + while True: + try: + ep = await ucp.create_endpoint(ucp.get_address(), remote_port) + return ep + except ucp.exceptions.UCXCanceled as e: + print( + "%s[%d]->%s[%d] Failed: %s" + % (my_task, my_port, remote_task, remote_port, e), + flush=True, + ) + await self._sleep(0.1) - if ep is None: - ep = await create_endpoint_retry(my_port, port, "Worker", "Worker") + def get_connections(self): + return self.connections - t = monotonic() - if GatherAsync: - msgs = [ep.recv(msg2recv), ep.send(msg2send)] * Iterations - await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) - else: - for i in range(Iterations): - msgs = [ep.recv(msg2recv), ep.send(msg2send)] - await asyncio.gather(*msgs, loop=asyncio.get_event_loop()) - if cache_only is False: - bytes_bandwidth[remote_port].append( - (send_recv_bytes, send_recv_bytes / (monotonic() - t)) - ) + async def run(self): + self._init() # Start listener - listener = ucp.create_listener(_listener_cb) - with lock: - signal[0] += 1 - ports[worker_num] = listener.port + listener = self._create_listener() + self.listener_port = listener.port - while signal[0] != num_workers: + with self.lock: + self.signal[0] += 1 + self.ports[self.worker_num] = listener.port + + while self.signal[0] != self.num_workers: pass if PersistentEndpoints: - for i in range(endpoints_per_worker): + for i in range(self.endpoints_per_worker): client_tasks = [] # Create endpoints to all other workers - for remote_port in list(ports): + for remote_port in list(self.ports): if remote_port == listener.port: continue - ep = await create_endpoint_retry( - listener.port, remote_port, "Worker", "Worker" - ) - bytes_bandwidth[remote_port] = [] - eps[(remote_port, i)] = ep + + ep = await self._create_endpoint(remote_port) + self.bytes_bandwidth[remote_port] = [] + self.conns[(remote_port, i)] = ep client_tasks.append( - _client(listener.port, remote_port, ep, cache_only=True) + self._client(listener.port, remote_port, ep, cache_only=True) ) - await asyncio.gather(*client_tasks, loop=asyncio.get_event_loop()) + await self._gather(client_tasks) - # Wait until listener_eps have all been cached - while len(listener_eps) != endpoints_per_worker * (num_workers - 1): - await asyncio.sleep(0.1) + # 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) # Exchange messages with other workers for i in range(3): client_tasks = [] listener_tasks = [] - for (remote_port, _), ep in eps.items(): - client_tasks.append(_client(listener.port, remote_port, ep)) - for listener_ep in listener_eps: - listener_tasks.append(_listener(listener_ep)) + for (remote_port, _), ep in self.conns.items(): + client_tasks.append(self._client(listener.port, remote_port, ep)) + for listener_ep in self.get_connections(): + listener_tasks.append(self._listener(listener_ep)) all_tasks = client_tasks + listener_tasks - await asyncio.gather(*all_tasks, loop=asyncio.get_event_loop()) + await self._gather(all_tasks) else: for i in range(3): # Create endpoints to all other workers client_tasks = [] - for port in list(ports): + for port in list(self.ports): if port == listener.port: continue - client_tasks.append(_client(listener.port, port)) - await asyncio.gather(*client_tasks, loop=asyncio.get_event_loop()) + client_tasks.append(self._client(listener.port, port)) + await self._gather(client_tasks) - with lock: - signal[1] += 1 - ports[worker_num] = listener.port + with self.lock: + self.signal[1] += 1 + self.ports[self.worker_num] = listener.port - while signal[1] != num_workers: + while self.signal[1] != self.num_workers: pass - for remote_port, bb in bytes_bandwidth.items(): + listener.close() + + # Wait for a shutdown signal from monitor + try: + while not listener.closed(): + await self._sleep(0.1) + except ucp.UCXCloseError: + pass + + def get_results(self): + for remote_port, bb in self.bytes_bandwidth.items(): 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)) @@ -154,7 +189,7 @@ async def _client(my_port, remote_port, ep=None, cache_only=False): "[%d, %d] Transferred bytes: %s, average bandwidth: %s/s, " "median bandwidth: %s/s" % ( - listener.port, + self.listener_port, remote_port, format_bytes(total_bytes), format_bytes(avg_bandwidth), @@ -162,17 +197,6 @@ async def _client(my_port, remote_port, ep=None, cache_only=False): ) ) - listener.close() - - # Wait for a shutdown signal from monitor - try: - while not listener.closed(): - await asyncio.sleep(0.1) - except ucp.UCXCloseError: - pass - - asyncio.get_event_loop().run_until_complete(_worker()) - def tornado_worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): conns = dict() @@ -217,9 +241,6 @@ async def _client(my_port, remote_port, conn=None): message_size, frames = await _get_message_size_and_frames() send_recv_bytes = (message_size * 2) * Iterations - # if ep is None: - # ep = await create_endpoint_retry(my_port, port, "Worker", "Worker") - t = monotonic() if GatherAsync: msgs = [ @@ -329,6 +350,12 @@ async def _client(my_port, remote_port, conn=None): IOLoop.current().run_sync(_worker) +def base_worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): + w = BaseWorker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker) + asyncio.get_event_loop().run_until_complete(w.run()) + w.get_results() + + def _test_send_recv_cu(num_workers, endpoints_per_worker): ctx = multiprocessing.get_context("spawn") @@ -340,9 +367,9 @@ def _test_send_recv_cu(num_workers, endpoints_per_worker): for worker_num in range(num_workers): worker_process = ctx.Process( name="worker", - # target=worker, + target=base_worker, # target=asyncio_worker, - target=tornado_worker, + # target=tornado_worker, args=[signal, ports, lock, worker_num, num_workers, endpoints_per_worker], ) worker_process.start() From 4c485c225c5ff5c952e35523048664df952855e5 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Tue, 6 Jul 2021 06:17:26 -0700 Subject: [PATCH 13/42] Use BaseWorker class to create TornadoWorker all-to-all benchmark --- tests/test_multiple_processes_all_to_all.py | 275 ++++++++------------ tests/utils.py | 14 + 2 files changed, 126 insertions(+), 163 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index 33dcd20fe..0fbc68a16 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -23,7 +23,14 @@ class BaseWorker: def __init__( - self, signal, ports, lock, worker_num, num_workers, endpoints_per_worker + self, + signal, + ports, + lock, + worker_num, + num_workers, + endpoints_per_worker, + transfer_to_cache, ): self.signal = signal self.ports = ports @@ -32,6 +39,10 @@ def __init__( self.num_workers = num_workers self.endpoints_per_worker = endpoints_per_worker + # UCX only effectively creates endpoints at first transfer, but this + # isn't necessary for tornado/asyncio. + self.transfer_to_cache = transfer_to_cache + self.conns = dict() self.connections = set() self.bytes_bandwidth = dict() @@ -44,6 +55,12 @@ async def _sleep(self, delay): async def _gather(self, tasks): await asyncio.gather(*tasks) + async def _get_messages_and_size(self): + msg2send = np.arange(Size) + msg2recv = np.empty_like(msg2send) + + return (msg2send, msg2recv, nbytes(msg2send)) + async def _transfer(self, ep, msg2send, msg2recv, send_first=True): if GatherAsync: msgs = [ep.send(msg2send), ep.recv(msg2recv)] * Iterations @@ -54,15 +71,13 @@ async def _transfer(self, ep, msg2send, msg2recv, send_first=True): await self._gather(msgs) async def _listener(self, ep): - msg2send = np.arange(Size) - msg2recv = np.empty_like(msg2send) + msg2send, msg2recv, _ = await self._get_messages_and_size() await self._transfer(ep, msg2send, msg2recv) async def _client(self, my_port, remote_port, ep=None, cache_only=False): - msg2send = np.arange(Size) - msg2recv = np.empty_like(msg2send) - send_recv_bytes = (nbytes(msg2send) + nbytes(msg2recv)) * Iterations + msg2send, msg2recv, msg_size = await self._get_messages_and_size() + send_recv_bytes = (msg_size * 2) * Iterations if ep is None: ep = await self._create_endpoint(remote_port) @@ -84,11 +99,11 @@ async def _listener_cb(self, ep): self.connections.add(ep) await self._listener(ep) - def _create_listener(self): + async def _create_listener(self): return ucp.create_listener(self._listener_cb) async def _create_endpoint(self, remote_port): - my_port = self.listener_port + my_port = self.listener.port my_task = "Worker" remote_task = "Worker" @@ -111,12 +126,11 @@ async def run(self): self._init() # Start listener - listener = self._create_listener() - self.listener_port = listener.port + self.listener = await self._create_listener() with self.lock: self.signal[0] += 1 - self.ports[self.worker_num] = listener.port + self.ports[self.worker_num] = self.listener.port while self.signal[0] != self.num_workers: pass @@ -126,16 +140,21 @@ async def run(self): client_tasks = [] # Create endpoints to all other workers for remote_port in list(self.ports): - if remote_port == listener.port: + if remote_port == self.listener.port: continue ep = await self._create_endpoint(remote_port) self.bytes_bandwidth[remote_port] = [] self.conns[(remote_port, i)] = ep - client_tasks.append( - self._client(listener.port, remote_port, ep, cache_only=True) - ) - await self._gather(client_tasks) + + if self.transfer_to_cache: + client_tasks.append( + self._client( + self.listener.port, remote_port, ep, cache_only=True + ) + ) + if self.transfer_to_cache: + await self._gather(client_tasks) # Wait until listener->ep connections have all been cached while len(self.get_connections()) != self.endpoints_per_worker * ( @@ -148,7 +167,9 @@ async def run(self): client_tasks = [] listener_tasks = [] for (remote_port, _), ep in self.conns.items(): - client_tasks.append(self._client(listener.port, remote_port, ep)) + client_tasks.append( + self._client(self.listener.port, remote_port, ep) + ) for listener_ep in self.get_connections(): listener_tasks.append(self._listener(listener_ep)) @@ -159,23 +180,26 @@ async def run(self): # Create endpoints to all other workers client_tasks = [] for port in list(self.ports): - if port == listener.port: + if port == self.listener.port: continue - client_tasks.append(self._client(listener.port, port)) + client_tasks.append(self._client(self.listener.port, port)) await self._gather(client_tasks) with self.lock: self.signal[1] += 1 - self.ports[self.worker_num] = listener.port + self.ports[self.worker_num] = self.listener.port while self.signal[1] != self.num_workers: pass - listener.close() + for conn in self.get_connections(): + await conn.close() + + self.listener.close() # Wait for a shutdown signal from monitor try: - while not listener.closed(): + while not self.listener.closed(): await self._sleep(0.1) except ucp.UCXCloseError: pass @@ -189,7 +213,7 @@ def get_results(self): "[%d, %d] Transferred bytes: %s, average bandwidth: %s/s, " "median bandwidth: %s/s" % ( - self.listener_port, + self.listener.port, remote_port, format_bytes(total_bytes), format_bytes(avg_bandwidth), @@ -198,164 +222,89 @@ def get_results(self): ) -def tornado_worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): - conns = dict() - bytes_bandwidth = dict() - - global cluster_started - cluster_started = False - - async def _worker(): - def _register_cluster_started(): - global cluster_started - cluster_started = True - - async def _get_message_size_and_frames(): - message = np.arange(Size) - - msg = {"data": to_serialize(message)} - frames = await to_frames(msg, serializers=("cuda", "dask", "pickle")) - - return (nbytes(message), frames) - - async def _listener(conn): - _, frames = await _get_message_size_and_frames() - - if GatherAsync: - msgs = [ - op - for i in range(Iterations) - for op in [conn.recv(), conn.send(frames)] - ] - await gen.multi(msgs) - else: - for i in range(Iterations): - msgs = [conn.send(frames), conn.recv()] - await gen.multi(msgs) - - # This seems to be faster! - # await conn.send(frames) - # await conn.recv() - - async def _client(my_port, remote_port, conn=None): - message_size, frames = await _get_message_size_and_frames() - send_recv_bytes = (message_size * 2) * Iterations - - t = monotonic() - if GatherAsync: - msgs = [ - op - for i in range(Iterations) - for op in [conn.recv(), conn.send(frames)] - ] - await gen.multi(msgs) - else: - for i in range(Iterations): - msgs = [conn.recv(), conn.send(frames)] - await gen.multi(msgs) - - # This seems to be faster! - # await conn.recv() - # await conn.send(frames) - - bytes_bandwidth[remote_port].append( - (send_recv_bytes, send_recv_bytes / (monotonic() - t)) - ) - - host = ucp.get_address(ifname="enp1s0f0") - - # Start listener - listener = await TornadoTCPServer.start_server(host, None) - with lock: - signal[0] += 1 - ports[worker_num] = listener.port +class TornadoWorker(BaseWorker): + def __init__( + self, signal, ports, lock, worker_num, num_workers, endpoints_per_worker + ): + super().__init__( + signal, + ports, + lock, + worker_num, + num_workers, + endpoints_per_worker, + transfer_to_cache=False, + ) - while signal[0] != num_workers: - await gen.sleep(0) + async def _sleep(self, delay): + await gen.sleep(delay) - print(list(ports)) + async def _gather(self, tasks): + await gen.multi(tasks) - if PersistentEndpoints: - for i in range(endpoints_per_worker): - client_tasks = [] - # Create endpoints to all other workers - for remote_port in list(ports): - if remote_port == listener.port: - continue - conn = await TornadoTCPConnection.connect(host, remote_port) - conns[(remote_port, i)] = conn - bytes_bandwidth[remote_port] = [] + async def _get_messages_and_size(self): + message = np.arange(Size) - # Wait until all clients connected to listener - while len(listener.get_connections()) != endpoints_per_worker * ( - num_workers - 1 - ): - await gen.sleep(0) + msg = {"data": to_serialize(message)} + frames = await to_frames(msg, serializers=("cuda", "dask", "pickle")) - # Exchange messages with other workers - for i in range(3): - client_tasks = [] - listener_tasks = [] - for (remote_port, _), conn in conns.items(): - client_tasks.append(_client(listener.port, remote_port, conn)) - for conn in listener.get_connections(): - listener_tasks.append(_listener(conn)) + return (frames, None, nbytes(message)) - all_tasks = client_tasks + listener_tasks - await gen.multi(all_tasks) + async def _transfer(self, ep, msg2send, msg2recv, send_first=True): + if GatherAsync: + msgs = [ + op for i in range(Iterations) for op in [ep.recv(), ep.send(msg2send)] + ] + await gen.multi(msgs) else: - for i in range(3): - # Create endpoints to all other workers - client_tasks = [] - for port in list(ports): - if port == listener.port: - continue - client_tasks.append(_client(listener.port, port)) - await gen.multi(client_tasks) + for i in range(Iterations): + msgs = [ep.recv(), ep.send(msg2send)] + await gen.multi(msgs) - with lock: - signal[1] += 1 - ports[worker_num] = listener.port + # This seems to be faster! + # await conn.recv() + # await conn.send(frames) - while signal[1] != num_workers: - pass + def _init(self): + return - for remote_port, bb in bytes_bandwidth.items(): - 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( - "[%d, %d] Transferred bytes: %s, average bandwidth: %s/s, " - "median bandwidth: %s/s" - % ( - listener.port, - remote_port, - format_bytes(total_bytes), - format_bytes(avg_bandwidth), - format_bytes(median_bandwidth), - ) - ) + async def _listener_cb(self, ep): + return - # listener.server.close() - for conn in listener.get_connections(): - conn.stream.close() + async def _create_listener(self): + host = ucp.get_address(ifname="enp1s0f0") + return await TornadoTCPServer.start_server(host, None) - # Wait for a shutdown signal from monitor - try: - while not all(c.stream.closed() for c in listener.get_connections()): - await gen.sleep(0) - except ucp.UCXCloseError: - pass + async def _create_endpoint(self, remote_port): + host = ucp.get_address(ifname="enp1s0f0") + return await TornadoTCPConnection.connect(host, remote_port) - IOLoop.current().run_sync(_worker) + def get_connections(self): + return self.listener._connections def base_worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): - w = BaseWorker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker) + w = BaseWorker( + signal, + ports, + lock, + worker_num, + num_workers, + endpoints_per_worker, + transfer_to_cache=True, + ) asyncio.get_event_loop().run_until_complete(w.run()) w.get_results() +def tornado_worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): + w = TornadoWorker( + signal, ports, lock, worker_num, num_workers, endpoints_per_worker + ) + IOLoop.current().run_sync(w.run) + w.get_results() + + def _test_send_recv_cu(num_workers, endpoints_per_worker): ctx = multiprocessing.get_context("spawn") @@ -367,9 +316,9 @@ def _test_send_recv_cu(num_workers, endpoints_per_worker): for worker_num in range(num_workers): worker_process = ctx.Process( name="worker", - target=base_worker, + # target=base_worker, # target=asyncio_worker, - # target=tornado_worker, + target=tornado_worker, args=[signal, ports, lock, worker_num, num_workers, endpoints_per_worker], ) worker_process.start() diff --git a/tests/utils.py b/tests/utils.py index 23aeea836..e72534f63 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -230,6 +230,13 @@ async def recv(self): raise Exception("aborted stream on truncated data") return 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): @@ -274,3 +281,10 @@ def get_connections(self): @property def port(self): return self._port + + def close(self): + self.server.stop() + self._closed = True + + def closed(self): + return self._closed From 6ddcf048d66b1bc3afb872bf398f06cc0f270421 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Tue, 6 Jul 2021 07:53:33 -0700 Subject: [PATCH 14/42] Add asyncio all-to-all benchmark --- tests/test_multiple_processes_all_to_all.py | 75 ++++++++++++++++++-- tests/utils.py | 77 +++++++++++++++++++++ 2 files changed, 147 insertions(+), 5 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index 0fbc68a16..746e8427c 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -6,7 +6,12 @@ import pytest from tornado import gen from tornado.ioloop import IOLoop -from utils import TornadoTCPConnection, TornadoTCPServer +from utils import ( + AsyncioCommConnection, + AsyncioCommServer, + TornadoTCPConnection, + TornadoTCPServer, +) from dask.utils import format_bytes from distributed.comm.utils import to_frames @@ -133,6 +138,7 @@ async def run(self): self.ports[self.worker_num] = self.listener.port while self.signal[0] != self.num_workers: + await self._sleep(0.1) pass if PersistentEndpoints: @@ -255,11 +261,11 @@ async def _transfer(self, ep, msg2send, msg2recv, send_first=True): msgs = [ op for i in range(Iterations) for op in [ep.recv(), ep.send(msg2send)] ] - await gen.multi(msgs) + await self._gather(msgs) else: for i in range(Iterations): msgs = [ep.recv(), ep.send(msg2send)] - await gen.multi(msgs) + await self._gather(msgs) # This seems to be faster! # await conn.recv() @@ -283,6 +289,57 @@ def get_connections(self): return self.listener._connections +class AsyncioWorker(BaseWorker): + def __init__( + self, signal, ports, lock, worker_num, num_workers, endpoints_per_worker + ): + super().__init__( + signal, + ports, + lock, + worker_num, + num_workers, + endpoints_per_worker, + transfer_to_cache=False, + ) + + async def _get_messages_and_size(self): + message = np.arange(Size) + + msg = {"data": to_serialize(message)} + frames = await to_frames(msg, serializers=("cuda", "dask", "pickle")) + + return (frames, None, nbytes(message)) + + async def _transfer(self, ep, msg2send, msg2recv, send_first=True): + if GatherAsync: + msgs = [ + op for i in range(Iterations) for op in [ep.recv(), ep.send(msg2send)] + ] + await self._gather(msgs) + else: + for i in range(Iterations): + msgs = [ep.recv(), ep.send(msg2send)] + await self._gather(msgs) + + def _init(self): + return + + async def _listener_cb(self, ep): + return + + async def _create_listener(self): + host = ucp.get_address(ifname="enp1s0f0") + return await AsyncioCommServer.start_server(host, None) + + async def _create_endpoint(self, remote_port): + host = ucp.get_address(ifname="enp1s0f0") + return await AsyncioCommConnection.open_connection(host, remote_port) + + def get_connections(self): + return self.listener._connections + + def base_worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): w = BaseWorker( signal, @@ -297,6 +354,14 @@ def base_worker(signal, ports, lock, worker_num, num_workers, endpoints_per_work w.get_results() +def asyncio_worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): + w = AsyncioWorker( + signal, ports, lock, worker_num, num_workers, endpoints_per_worker, + ) + asyncio.get_event_loop().run_until_complete(w.run()) + w.get_results() + + def tornado_worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): w = TornadoWorker( signal, ports, lock, worker_num, num_workers, endpoints_per_worker @@ -317,8 +382,8 @@ def _test_send_recv_cu(num_workers, endpoints_per_worker): worker_process = ctx.Process( name="worker", # target=base_worker, - # target=asyncio_worker, - target=tornado_worker, + target=asyncio_worker, + # target=tornado_worker, args=[signal, ports, lock, worker_num, num_workers, endpoints_per_worker], ) worker_process.start() diff --git a/tests/utils.py b/tests/utils.py index e72534f63..85636d8a6 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -1,3 +1,5 @@ +import asyncio +import functools import io import logging import os @@ -149,6 +151,7 @@ class TornadoTCPConnection: def __init__(self, stream, client=None): self._client = client self.stream = stream + self._closed = False @classmethod async def connect(cls, host, port): @@ -288,3 +291,77 @@ def close(self): 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, frames): + 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): + def _server_callback(connections, reader, writer): + connections.append(AsyncioCommConnection(reader, writer)) + + connections = [] + + server = await asyncio.start_server( + functools.partial(_server_callback, connections), 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() From 9675844ac3cc4765b7d7f4755e8cfdb20cd71faf Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Tue, 6 Jul 2021 13:19:26 -0700 Subject: [PATCH 15/42] Remove non-persistent endpoints --- tests/test_multiple_processes_all_to_all.py | 90 +++++++++------------ 1 file changed, 37 insertions(+), 53 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index 746e8427c..335376681 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -20,7 +20,6 @@ import ucp -PersistentEndpoints = True GatherAsync = False Iterations = 3 Size = 2 ** 15 @@ -61,7 +60,7 @@ async def _gather(self, tasks): await asyncio.gather(*tasks) async def _get_messages_and_size(self): - msg2send = np.arange(Size) + msg2send = np.arange(Size, dtype=np.uint8) msg2recv = np.empty_like(msg2send) return (msg2send, msg2recv, nbytes(msg2send)) @@ -100,8 +99,7 @@ def _init(self): ucp.init() async def _listener_cb(self, ep): - if PersistentEndpoints: - self.connections.add(ep) + self.connections.add(ep) await self._listener(ep) async def _create_listener(self): @@ -139,58 +137,44 @@ async def run(self): while self.signal[0] != self.num_workers: await self._sleep(0.1) - pass - if PersistentEndpoints: - for i in range(self.endpoints_per_worker): - client_tasks = [] - # Create endpoints to all other workers - for remote_port in list(self.ports): - if remote_port == self.listener.port: - continue - - ep = await self._create_endpoint(remote_port) - self.bytes_bandwidth[remote_port] = [] - self.conns[(remote_port, i)] = ep - - if self.transfer_to_cache: - client_tasks.append( - self._client( - self.listener.port, remote_port, ep, cache_only=True - ) - ) - if self.transfer_to_cache: - await self._gather(client_tasks) + for i in range(self.endpoints_per_worker): + client_tasks = [] + # Create endpoints to all other workers + for remote_port in list(self.ports): + if remote_port == self.listener.port: + continue - # 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) + ep = await self._create_endpoint(remote_port) + self.bytes_bandwidth[remote_port] = [] + self.conns[(remote_port, i)] = ep - # Exchange messages with other workers - for i in range(3): - client_tasks = [] - listener_tasks = [] - for (remote_port, _), ep in self.conns.items(): + if self.transfer_to_cache: client_tasks.append( - self._client(self.listener.port, remote_port, ep) + self._client( + self.listener.port, remote_port, ep, cache_only=True + ) ) - 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) - else: - for i in range(3): - # Create endpoints to all other workers - client_tasks = [] - for port in list(self.ports): - if port == self.listener.port: - continue - client_tasks.append(self._client(self.listener.port, port)) + if self.transfer_to_cache: await self._gather(client_tasks) + # 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) + + # Exchange messages with other workers + client_tasks = [] + listener_tasks = [] + for (remote_port, _), ep in self.conns.items(): + client_tasks.append(self._client(self.listener.port, remote_port, 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) + with self.lock: self.signal[1] += 1 self.ports[self.worker_num] = self.listener.port @@ -249,7 +233,7 @@ async def _gather(self, tasks): await gen.multi(tasks) async def _get_messages_and_size(self): - message = np.arange(Size) + message = np.arange(Size, dtype=np.uint8) msg = {"data": to_serialize(message)} frames = await to_frames(msg, serializers=("cuda", "dask", "pickle")) @@ -304,7 +288,7 @@ def __init__( ) async def _get_messages_and_size(self): - message = np.arange(Size) + message = np.arange(Size, dtype=np.uint8) msg = {"data": to_serialize(message)} frames = await to_frames(msg, serializers=("cuda", "dask", "pickle")) @@ -382,8 +366,8 @@ def _test_send_recv_cu(num_workers, endpoints_per_worker): worker_process = ctx.Process( name="worker", # target=base_worker, - target=asyncio_worker, - # target=tornado_worker, + # target=asyncio_worker, + target=tornado_worker, args=[signal, ports, lock, worker_num, num_workers, endpoints_per_worker], ) worker_process.start() From c8fb2a7a00a26600ff434d266a7bffac1c323bba Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Wed, 7 Jul 2021 07:01:54 -0700 Subject: [PATCH 16/42] Split UCX transfers in frames to make it similar to Tornado/Asyncio --- tests/test_multiple_processes_all_to_all.py | 104 ++++++-------------- tests/utils.py | 87 ++++++++++++++++ 2 files changed, 119 insertions(+), 72 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index 335376681..17444370f 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -11,6 +11,8 @@ AsyncioCommServer, TornadoTCPConnection, TornadoTCPServer, + UCXConnection, + UCXServer, ) from dask.utils import format_bytes @@ -22,7 +24,7 @@ GatherAsync = False Iterations = 3 -Size = 2 ** 15 +Size = 2 ** 10 class BaseWorker: @@ -60,19 +62,32 @@ async def _gather(self, tasks): await asyncio.gather(*tasks) async def _get_messages_and_size(self): - msg2send = np.arange(Size, dtype=np.uint8) - msg2recv = np.empty_like(msg2send) + message = np.arange(Size, dtype=np.uint8) + + msg = {"data": to_serialize(message)} + frames = await to_frames(msg, serializers=("cuda", "dask", "pickle")) - return (msg2send, msg2recv, nbytes(msg2send)) + return (frames, None, nbytes(message)) async def _transfer(self, ep, msg2send, msg2recv, send_first=True): if GatherAsync: - msgs = [ep.send(msg2send), ep.recv(msg2recv)] * Iterations + msgs = [ + op for i in range(Iterations) for op in [ep.recv(), ep.send(msg2send)] + ] await self._gather(msgs) else: for i in range(Iterations): - msgs = [ep.send(msg2send), ep.recv(msg2recv)] - await self._gather(msgs) + if False: + 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) async def _listener(self, ep): msg2send, msg2recv, _ = await self._get_messages_and_size() @@ -98,12 +113,9 @@ async def _client(self, my_port, remote_port, ep=None, cache_only=False): def _init(self): ucp.init() - async def _listener_cb(self, ep): - self.connections.add(ep) - await self._listener(ep) - async def _create_listener(self): - return ucp.create_listener(self._listener_cb) + host = ucp.get_address(ifname="enp1s0f0") + return await UCXServer.start_server(host, 0, self._listener) async def _create_endpoint(self, remote_port): my_port = self.listener.port @@ -112,7 +124,7 @@ async def _create_endpoint(self, remote_port): while True: try: - ep = await ucp.create_endpoint(ucp.get_address(), remote_port) + ep = await UCXConnection.open_connection(ucp.get_address(), remote_port) return ep except ucp.exceptions.UCXCanceled as e: print( @@ -123,7 +135,7 @@ async def _create_endpoint(self, remote_port): await self._sleep(0.1) def get_connections(self): - return self.connections + return self.listener._connections async def run(self): self._init() @@ -155,6 +167,8 @@ async def run(self): self.listener.port, remote_port, ep, cache_only=True ) ) + # for listener_ep in self.get_connections(): + # client_tasks.append(self._listener(listener_ep)) if self.transfer_to_cache: await self._gather(client_tasks) @@ -171,7 +185,6 @@ async def run(self): client_tasks.append(self._client(self.listener.port, remote_port, 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) @@ -180,7 +193,7 @@ async def run(self): self.ports[self.worker_num] = self.listener.port while self.signal[1] != self.num_workers: - pass + await self._sleep(0) for conn in self.get_connections(): await conn.close() @@ -196,6 +209,7 @@ async def run(self): def get_results(self): for remote_port, bb in self.bytes_bandwidth.items(): + print(bb) 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)) @@ -232,35 +246,9 @@ async def _sleep(self, delay): async def _gather(self, tasks): await gen.multi(tasks) - async def _get_messages_and_size(self): - message = np.arange(Size, dtype=np.uint8) - - msg = {"data": to_serialize(message)} - frames = await to_frames(msg, serializers=("cuda", "dask", "pickle")) - - return (frames, None, nbytes(message)) - - async def _transfer(self, ep, msg2send, msg2recv, send_first=True): - if GatherAsync: - msgs = [ - op for i in range(Iterations) for op in [ep.recv(), ep.send(msg2send)] - ] - await self._gather(msgs) - else: - for i in range(Iterations): - msgs = [ep.recv(), ep.send(msg2send)] - await self._gather(msgs) - - # This seems to be faster! - # await conn.recv() - # await conn.send(frames) - def _init(self): return - async def _listener_cb(self, ep): - return - async def _create_listener(self): host = ucp.get_address(ifname="enp1s0f0") return await TornadoTCPServer.start_server(host, None) @@ -269,9 +257,6 @@ async def _create_endpoint(self, remote_port): host = ucp.get_address(ifname="enp1s0f0") return await TornadoTCPConnection.connect(host, remote_port) - def get_connections(self): - return self.listener._connections - class AsyncioWorker(BaseWorker): def __init__( @@ -287,31 +272,9 @@ def __init__( transfer_to_cache=False, ) - async def _get_messages_and_size(self): - message = np.arange(Size, dtype=np.uint8) - - msg = {"data": to_serialize(message)} - frames = await to_frames(msg, serializers=("cuda", "dask", "pickle")) - - return (frames, None, nbytes(message)) - - async def _transfer(self, ep, msg2send, msg2recv, send_first=True): - if GatherAsync: - msgs = [ - op for i in range(Iterations) for op in [ep.recv(), ep.send(msg2send)] - ] - await self._gather(msgs) - else: - for i in range(Iterations): - msgs = [ep.recv(), ep.send(msg2send)] - await self._gather(msgs) - def _init(self): return - async def _listener_cb(self, ep): - return - async def _create_listener(self): host = ucp.get_address(ifname="enp1s0f0") return await AsyncioCommServer.start_server(host, None) @@ -320,9 +283,6 @@ async def _create_endpoint(self, remote_port): host = ucp.get_address(ifname="enp1s0f0") return await AsyncioCommConnection.open_connection(host, remote_port) - def get_connections(self): - return self.listener._connections - def base_worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): w = BaseWorker( @@ -365,9 +325,9 @@ def _test_send_recv_cu(num_workers, endpoints_per_worker): for worker_num in range(num_workers): worker_process = ctx.Process( name="worker", - # target=base_worker, + target=base_worker, # target=asyncio_worker, - target=tornado_worker, + # target=tornado_worker, args=[signal, ports, lock, worker_num, num_workers, endpoints_per_worker], ) worker_process.start() diff --git a/tests/utils.py b/tests/utils.py index 85636d8a6..8195e256e 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -306,8 +306,10 @@ async def open_connection(cls, host, port): async def send(self, frames): nframes = len(frames) self.writer.write(struct.pack("Q", nframes)) + # await self.writer.drain() sizes = list(nbytes(f) for f in frames) self.writer.write(struct.pack(nframes * "Q", *sizes)) + # await self.writer.drain() for f in frames: self.writer.write(f) await self.writer.drain() @@ -365,3 +367,88 @@ def close(self): 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, frames): + nframes = len(frames) + print(f"Send nframes: {nframes}") + 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)) + print(f"Send sizes: {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) + print(f"Recv nframes: {nframes}") + + sizes = np.empty((struct.calcsize(nframes[0] * "Q"),), dtype="u1") + await self.ep.recv(sizes) + sizes = struct.unpack(nframes[0] * "Q", sizes) + print(f"Recv sizes: {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): + async def _server_callback(connections, ep): + conn = UCXConnection(ep) + connections.append(UCXConnection(ep)) + await listener_func(conn) + + connections = [] + + server = ucp.create_listener( + functools.partial(_server_callback, connections), 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() From bf48ebd7ecbef67175b8b76c45a4c198021ad4b3 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Wed, 7 Jul 2021 14:16:08 -0700 Subject: [PATCH 17/42] Move serialization to comms send method --- tests/test_multiple_processes_all_to_all.py | 29 +++++++-------------- tests/utils.py | 18 ++++++++++--- 2 files changed, 23 insertions(+), 24 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index 17444370f..bfd44b160 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -16,15 +16,13 @@ ) from dask.utils import format_bytes -from distributed.comm.utils import to_frames -from distributed.protocol import to_serialize from distributed.utils import nbytes import ucp GatherAsync = False Iterations = 3 -Size = 2 ** 10 +Size = 2 ** 20 class BaseWorker: @@ -61,15 +59,7 @@ async def _sleep(self, delay): async def _gather(self, tasks): await asyncio.gather(*tasks) - async def _get_messages_and_size(self): - message = np.arange(Size, dtype=np.uint8) - - msg = {"data": to_serialize(message)} - frames = await to_frames(msg, serializers=("cuda", "dask", "pickle")) - - return (frames, None, nbytes(message)) - - async def _transfer(self, ep, msg2send, msg2recv, send_first=True): + async def _transfer(self, ep, msg2send, send_first=True): if GatherAsync: msgs = [ op for i in range(Iterations) for op in [ep.recv(), ep.send(msg2send)] @@ -90,19 +80,19 @@ async def _transfer(self, ep, msg2send, msg2recv, send_first=True): await ep.send(msg2send) async def _listener(self, ep): - msg2send, msg2recv, _ = await self._get_messages_and_size() + message = np.arange(Size, dtype=np.uint8) - await self._transfer(ep, msg2send, msg2recv) + await self._transfer(ep, message) async def _client(self, my_port, remote_port, ep=None, cache_only=False): - msg2send, msg2recv, msg_size = await self._get_messages_and_size() - send_recv_bytes = (msg_size * 2) * Iterations + message = np.arange(Size, dtype=np.uint8) + send_recv_bytes = (nbytes(message) * 2) * Iterations if ep is None: ep = await self._create_endpoint(remote_port) t = monotonic() - await self._transfer(ep, msg2send, msg2recv, send_first=False) + await self._transfer(ep, message, send_first=False) total_time = monotonic() - t if cache_only is False: @@ -209,7 +199,6 @@ async def run(self): def get_results(self): for remote_port, bb in self.bytes_bandwidth.items(): - print(bb) 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)) @@ -325,9 +314,9 @@ def _test_send_recv_cu(num_workers, endpoints_per_worker): for worker_num in range(num_workers): worker_process = ctx.Process( name="worker", - target=base_worker, + # target=base_worker, # target=asyncio_worker, - # target=tornado_worker, + target=tornado_worker, args=[signal, ports, lock, worker_num, num_workers, endpoints_per_worker], ) worker_process.start() diff --git a/tests/utils.py b/tests/utils.py index 8195e256e..bbf0aa409 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -11,7 +11,8 @@ from tornado.tcpclient import TCPClient from tornado.tcpserver import TCPServer -from distributed.comm.utils import from_frames +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 @@ -160,11 +161,14 @@ async def connect(cls, host, port): stream.set_nodelay(True) return cls(stream, client=client) - async def send(self, frames): + 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) @@ -303,7 +307,10 @@ 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, frames): + 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)) # await self.writer.drain() @@ -378,7 +385,10 @@ async def open_connection(cls, host, port): ep = await ucp.create_endpoint(host, port) return cls(ep) - async def send(self, frames): + async def send(self, message): + msg = {"data": to_serialize(message)} + frames = await to_frames(msg, serializers=("cuda", "dask", "pickle")) + nframes = len(frames) print(f"Send nframes: {nframes}") await self.ep.send(struct.pack("Q", nframes)) From 7d2eaff21da0b3a5d8a83b214313b81823e8a772 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Thu, 8 Jul 2021 06:02:55 -0700 Subject: [PATCH 18/42] Add monitor process to allow syncing workers over network --- tests/test_multiple_processes_all_to_all.py | 322 +++++++++++++++----- 1 file changed, 254 insertions(+), 68 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index bfd44b160..b7a45960b 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -24,6 +24,12 @@ Iterations = 3 Size = 2 ** 20 +OP_NONE = 0 +OP_WORKER_LISTENING = 1 +OP_CLUSTER_READY = 2 +OP_WORKER_COMPLETED = 3 +OP_SHUTDOWN = 4 + class BaseWorker: def __init__( @@ -34,6 +40,7 @@ def __init__( worker_num, num_workers, endpoints_per_worker, + monitor_port, transfer_to_cache, ): self.signal = signal @@ -47,6 +54,8 @@ def __init__( # isn't necessary for tornado/asyncio. self.transfer_to_cache = transfer_to_cache + self.monitor_port = monitor_port + self.conns = dict() self.connections = set() self.bytes_bandwidth = dict() @@ -74,7 +83,7 @@ async def _transfer(self, ep, msg2send, send_first=True): # This seems to be faster! if send_first: await ep.send(msg2send) - await ep.recv() + _, msg = await ep.recv() else: await ep.recv() await ep.send(msg2send) @@ -84,42 +93,44 @@ async def _listener(self, ep): await self._transfer(ep, message) - async def _client(self, my_port, remote_port, ep=None, cache_only=False): + async def _client(self, my_port, worker_address, ep, cache_only=False): message = np.arange(Size, dtype=np.uint8) send_recv_bytes = (nbytes(message) * 2) * Iterations - if ep is None: - ep = await self._create_endpoint(remote_port) - t = monotonic() await self._transfer(ep, message, send_first=False) total_time = monotonic() - t if cache_only is False: - self.bytes_bandwidth[remote_port].append( + self.bytes_bandwidth[worker_address].append( (send_recv_bytes, send_recv_bytes / total_time) ) def _init(self): ucp.init() - async def _create_listener(self): - host = ucp.get_address(ifname="enp1s0f0") - return await UCXServer.start_server(host, 0, self._listener) + async def _create_listener(self, host, port=None): + return await UCXServer.start_server(host, port, self._listener) - async def _create_endpoint(self, remote_port): + async def _create_monitor_listener(self, host, port=None): + async def _cb(ep): + pass + + return await UCXServer.start_server(host, port, _cb) + + 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(ucp.get_address(), remote_port) + 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, remote_port, e), + % (my_task, my_port, remote_task, port, e), flush=True, ) await self._sleep(0.1) @@ -127,38 +138,121 @@ async def _create_endpoint(self, remote_port): def get_connections(self): return self.listener._connections - async def run(self): + async def run_monitor(self): self._init() - # Start listener - self.listener = await self._create_listener() + self.listener_address = ucp.get_address(ifname="enp1s0f0") + self.listener = await self._create_monitor_listener( + self.listener_address, self.monitor_port + ) with self.lock: - self.signal[0] += 1 - self.ports[self.worker_num] = self.listener.port + self.signal[0] = self.listener.port - while self.signal[0] != self.num_workers: + # Wait for all workers to connect + while len(self.get_connections()) != ( + self.endpoints_per_worker * self.num_workers + ): await self._sleep(0.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,)) + + 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(self): + self._init() + + # Start listener + self.listener_address = ucp.get_address(ifname="enp1s0f0") + self.listener = await self._create_listener(self.listener_address) + + if self.monitor_port == 0: + with self.lock: + self.signal[0] += 1 + self.ports[self.worker_num] = self.listener.port + + while self.signal[0] != self.num_workers: + await self._sleep(0.1) + else: + monitor_ep = await self._create_endpoint( + self.listener_address, self.monitor_port + ) + await monitor_ep.send( + (OP_WORKER_LISTENING, self.listener_address, self.listener.port) + ) + _, worker_addresses = await monitor_ep.recv() + assert worker_addresses["data"][0] == OP_CLUSTER_READY + worker_addresses = worker_addresses["data"][1] + assert len(worker_addresses) == self.num_workers + for i in range(self.endpoints_per_worker): client_tasks = [] # Create endpoints to all other workers - for remote_port in list(self.ports): - if remote_port == self.listener.port: - continue - - ep = await self._create_endpoint(remote_port) - self.bytes_bandwidth[remote_port] = [] - self.conns[(remote_port, i)] = ep - - if self.transfer_to_cache: - client_tasks.append( - self._client( - self.listener.port, remote_port, ep, cache_only=True + if self.monitor_port == 0: + for remote_port in list(self.ports): + if remote_port == self.listener.port: + continue + + ep = await self._create_endpoint(self.listener_address, remote_port) + self.bytes_bandwidth[remote_port] = [] + self.conns[(remote_port, i)] = ep + + if self.transfer_to_cache: + client_tasks.append( + self._client( + self.listener.port, remote_port, ep, cache_only=True + ) + ) + else: + for worker_address in 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 + ) ) - ) - # for listener_ep in self.get_connections(): - # client_tasks.append(self._listener(listener_ep)) + if self.transfer_to_cache: await self._gather(client_tasks) @@ -171,19 +265,25 @@ async def run(self): # Exchange messages with other workers client_tasks = [] listener_tasks = [] - for (remote_port, _), ep in self.conns.items(): - client_tasks.append(self._client(self.listener.port, remote_port, ep)) + 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) - with self.lock: - self.signal[1] += 1 - self.ports[self.worker_num] = self.listener.port + if self.monitor_port == 0: + with self.lock: + self.signal[1] += 1 + self.ports[self.worker_num] = self.listener.port - while self.signal[1] != self.num_workers: - await self._sleep(0) + while self.signal[1] != self.num_workers: + await self._sleep(0) + else: + await monitor_ep.send((OP_WORKER_COMPLETED,)) + _, shutdown = await monitor_ep.recv() + shutdown = shutdown["data"] + assert shutdown[0] == OP_SHUTDOWN for conn in self.get_connections(): await conn.close() @@ -203,7 +303,7 @@ def get_results(self): avg_bandwidth = np.mean(list(b[1] for b in bb)) median_bandwidth = np.median(list(b[1] for b in bb)) print( - "[%d, %d] Transferred bytes: %s, average bandwidth: %s/s, " + "[%d, %s] Transferred bytes: %s, average bandwidth: %s/s, " "median bandwidth: %s/s" % ( self.listener.port, @@ -217,7 +317,14 @@ def get_results(self): class TornadoWorker(BaseWorker): def __init__( - self, signal, ports, lock, worker_num, num_workers, endpoints_per_worker + self, + signal, + ports, + lock, + worker_num, + num_workers, + endpoints_per_worker, + monitor_port, ): super().__init__( signal, @@ -226,6 +333,7 @@ def __init__( worker_num, num_workers, endpoints_per_worker, + monitor_port, transfer_to_cache=False, ) @@ -238,18 +346,23 @@ async def _gather(self, tasks): def _init(self): return - async def _create_listener(self): - host = ucp.get_address(ifname="enp1s0f0") - return await TornadoTCPServer.start_server(host, None) + async def _create_listener(self, host, port=None): + return await TornadoTCPServer.start_server(host, port=None) - async def _create_endpoint(self, remote_port): - host = ucp.get_address(ifname="enp1s0f0") - return await TornadoTCPConnection.connect(host, remote_port) + async def _create_endpoint(self, host, port): + return await TornadoTCPConnection.connect(host, port) class AsyncioWorker(BaseWorker): def __init__( - self, signal, ports, lock, worker_num, num_workers, endpoints_per_worker + self, + signal, + ports, + lock, + worker_num, + num_workers, + endpoints_per_worker, + monitor_port, ): super().__init__( signal, @@ -258,22 +371,31 @@ def __init__( worker_num, num_workers, endpoints_per_worker, + monitor_port, transfer_to_cache=False, ) def _init(self): return - async def _create_listener(self): - host = ucp.get_address(ifname="enp1s0f0") - return await AsyncioCommServer.start_server(host, None) + async def _create_listener(self, host, port=None): + return await AsyncioCommServer.start_server(host, port) - async def _create_endpoint(self, remote_port): + async def _create_endpoint(self, host, port): host = ucp.get_address(ifname="enp1s0f0") - return await AsyncioCommConnection.open_connection(host, remote_port) - - -def base_worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): + return await AsyncioCommConnection.open_connection(host, port) + + +def base_worker( + signal, + ports, + lock, + worker_num, + num_workers, + endpoints_per_worker, + is_monitor, + monitor_port, +): w = BaseWorker( signal, ports, @@ -281,43 +403,104 @@ def base_worker(signal, ports, lock, worker_num, num_workers, endpoints_per_work worker_num, num_workers, endpoints_per_worker, + monitor_port, transfer_to_cache=True, ) - asyncio.get_event_loop().run_until_complete(w.run()) + run_func = w.run_monitor if is_monitor else w.run + asyncio.get_event_loop().run_until_complete(run_func()) w.get_results() -def asyncio_worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): +def asyncio_worker( + signal, + ports, + lock, + worker_num, + num_workers, + endpoints_per_worker, + is_monitor, + monitor_port, +): w = AsyncioWorker( - signal, ports, lock, worker_num, num_workers, endpoints_per_worker, + signal, + ports, + lock, + worker_num, + num_workers, + endpoints_per_worker, + monitor_port, ) - asyncio.get_event_loop().run_until_complete(w.run()) + run_func = w.run_monitor if is_monitor else w.run + asyncio.get_event_loop().run_until_complete(run_func()) w.get_results() -def tornado_worker(signal, ports, lock, worker_num, num_workers, endpoints_per_worker): +def tornado_worker( + signal, + ports, + lock, + worker_num, + num_workers, + endpoints_per_worker, + is_monitor, + monitor_port, +): w = TornadoWorker( - signal, ports, lock, worker_num, num_workers, endpoints_per_worker + signal, + ports, + lock, + worker_num, + num_workers, + endpoints_per_worker, + monitor_port, ) - IOLoop.current().run_sync(w.run) + run_func = w.run_monitor if is_monitor else w.run + IOLoop.current().run_sync(run_func) w.get_results() def _test_send_recv_cu(num_workers, endpoints_per_worker): ctx = multiprocessing.get_context("spawn") + enable_monitor = False + monitor_port = 0 + signal = ctx.Array("i", [0, 0]) ports = ctx.Array("i", range(num_workers)) lock = ctx.Lock() + comm_type = base_worker + # comm_type = asyncio_worker + # comm_type = tornado_worker + + if enable_monitor: + monitor_process = ctx.Process( + name="worker", + target=comm_type, + args=[signal, ports, lock, 0, num_workers, endpoints_per_worker, True, 0], + ) + 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=base_worker, - # target=asyncio_worker, - target=tornado_worker, - args=[signal, ports, lock, worker_num, num_workers, endpoints_per_worker], + target=comm_type, + args=[ + signal, + ports, + lock, + worker_num, + num_workers, + endpoints_per_worker, + False, + monitor_port, + ], ) worker_process.start() worker_processes.append(worker_process) @@ -325,6 +508,9 @@ def _test_send_recv_cu(num_workers, endpoints_per_worker): for worker_process in worker_processes: worker_process.join() + if enable_monitor: + monitor_process.join() + assert worker_process.exitcode == 0 From 8ae422a6d082096b7dfbcc27995255fdd45b6591 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Thu, 8 Jul 2021 06:23:13 -0700 Subject: [PATCH 19/42] Split `run` into multiple methods --- tests/test_multiple_processes_all_to_all.py | 70 +++++++++++++-------- 1 file changed, 44 insertions(+), 26 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index b7a45960b..d150cdbc0 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -190,13 +190,7 @@ async def run_monitor(self): except ucp.UCXCloseError: pass - async def run(self): - self._init() - - # Start listener - self.listener_address = ucp.get_address(ifname="enp1s0f0") - self.listener = await self._create_listener(self.listener_address) - + async def _wait_for_workers(self): if self.monitor_port == 0: with self.lock: self.signal[0] += 1 @@ -205,17 +199,32 @@ async def run(self): while self.signal[0] != self.num_workers: await self._sleep(0.1) else: - monitor_ep = await self._create_endpoint( + self.monitor_ep = await self._create_endpoint( self.listener_address, self.monitor_port ) - await monitor_ep.send( + await self.monitor_ep.send( (OP_WORKER_LISTENING, self.listener_address, self.listener.port) ) - _, worker_addresses = await monitor_ep.recv() + _, worker_addresses = await self.monitor_ep.recv() assert worker_addresses["data"][0] == OP_CLUSTER_READY - worker_addresses = worker_addresses["data"][1] - assert len(worker_addresses) == self.num_workers + self.worker_addresses = worker_addresses["data"][1] + assert len(self.worker_addresses) == self.num_workers + async def _wait_for_completion(self): + if self.monitor_port == 0: + with self.lock: + self.signal[1] += 1 + self.ports[self.worker_num] = self.listener.port + + 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 @@ -235,7 +244,7 @@ async def run(self): ) ) else: - for worker_address in worker_addresses: + for worker_address in self.worker_addresses: if ( worker_address[0] == self.listener_address and worker_address[1] == self.listener.port @@ -256,12 +265,14 @@ async def run(self): 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 = [] @@ -272,19 +283,7 @@ async def run(self): all_tasks = client_tasks + listener_tasks await self._gather(all_tasks) - if self.monitor_port == 0: - with self.lock: - self.signal[1] += 1 - self.ports[self.worker_num] = self.listener.port - - while self.signal[1] != self.num_workers: - await self._sleep(0) - else: - await monitor_ep.send((OP_WORKER_COMPLETED,)) - _, shutdown = await monitor_ep.recv() - shutdown = shutdown["data"] - assert shutdown[0] == OP_SHUTDOWN - + async def _close_connections_and_listener(self): for conn in self.get_connections(): await conn.close() @@ -297,6 +296,25 @@ async def run(self): except ucp.UCXCloseError: pass + async def run(self): + self._init() + + # Start listener + self.listener_address = ucp.get_address(ifname="enp1s0f0") + 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() + def get_results(self): for remote_port, bb in self.bytes_bandwidth.items(): total_bytes = sum(b[0] for b in bb) From da2320e7ec1858f0d33fea2d76c234d5bc77403d Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Thu, 8 Jul 2021 09:36:07 -0700 Subject: [PATCH 20/42] Parametrize communication and monitor enable/disable --- tests/test_multiple_processes_all_to_all.py | 90 ++++++++++----------- tests/utils.py | 10 +-- 2 files changed, 43 insertions(+), 57 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index d150cdbc0..f84e4f6b3 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -20,7 +20,7 @@ import ucp -GatherAsync = False +GatherSendRecv = False Iterations = 3 Size = 2 ** 20 @@ -31,7 +31,7 @@ OP_SHUTDOWN = 4 -class BaseWorker: +class UCXProcess: def __init__( self, signal, @@ -69,24 +69,18 @@ async def _gather(self, tasks): await asyncio.gather(*tasks) async def _transfer(self, ep, msg2send, send_first=True): - if GatherAsync: - msgs = [ - op for i in range(Iterations) for op in [ep.recv(), ep.send(msg2send)] - ] - await self._gather(msgs) - else: - for i in range(Iterations): - if False: - msgs = [ep.recv(), ep.send(msg2send)] - await self._gather(msgs) + for i in range(Iterations): + if GatherSendRecv: + 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: - # This seems to be faster! - if send_first: - await ep.send(msg2send) - _, msg = await ep.recv() - else: - await ep.recv() - await ep.send(msg2send) + await ep.recv() + await ep.send(msg2send) async def _listener(self, ep): message = np.arange(Size, dtype=np.uint8) @@ -109,14 +103,11 @@ async def _client(self, my_port, worker_address, ep, cache_only=False): def _init(self): ucp.init() - async def _create_listener(self, host, port=None): - return await UCXServer.start_server(host, port, self._listener) - - async def _create_monitor_listener(self, host, port=None): - async def _cb(ep): - pass + async def _create_listener(self, host, port=None, cb=None): + return await UCXServer.start_server(host, port, cb or self._listener) - return await UCXServer.start_server(host, port, _cb) + async def _monitor_listener_cb(self, ep): + pass async def _create_endpoint(self, host, port): my_port = self.listener.port @@ -142,8 +133,8 @@ async def run_monitor(self): self._init() self.listener_address = ucp.get_address(ifname="enp1s0f0") - self.listener = await self._create_monitor_listener( - self.listener_address, self.monitor_port + self.listener = await self._create_listener( + self.listener_address, self.monitor_port, self._monitor_listener_cb, ) with self.lock: @@ -333,7 +324,7 @@ def get_results(self): ) -class TornadoWorker(BaseWorker): +class TornadoProcess(UCXProcess): def __init__( self, signal, @@ -364,14 +355,14 @@ async def _gather(self, tasks): def _init(self): return - async def _create_listener(self, host, port=None): - return await TornadoTCPServer.start_server(host, port=None) + 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 AsyncioWorker(BaseWorker): +class AsyncioProcess(UCXProcess): def __init__( self, signal, @@ -396,7 +387,7 @@ def __init__( def _init(self): return - async def _create_listener(self, host, port=None): + async def _create_listener(self, host, port=None, cb=None): return await AsyncioCommServer.start_server(host, port) async def _create_endpoint(self, host, port): @@ -404,7 +395,7 @@ async def _create_endpoint(self, host, port): return await AsyncioCommConnection.open_connection(host, port) -def base_worker( +def ucx_process( signal, ports, lock, @@ -414,7 +405,7 @@ def base_worker( is_monitor, monitor_port, ): - w = BaseWorker( + w = UCXProcess( signal, ports, lock, @@ -429,7 +420,7 @@ def base_worker( w.get_results() -def asyncio_worker( +def asyncio_process( signal, ports, lock, @@ -439,7 +430,7 @@ def asyncio_worker( is_monitor, monitor_port, ): - w = AsyncioWorker( + w = AsyncioProcess( signal, ports, lock, @@ -453,7 +444,7 @@ def asyncio_worker( w.get_results() -def tornado_worker( +def tornado_process( signal, ports, lock, @@ -463,7 +454,7 @@ def tornado_worker( is_monitor, monitor_port, ): - w = TornadoWorker( + w = TornadoProcess( signal, ports, lock, @@ -477,24 +468,21 @@ def tornado_worker( w.get_results() -def _test_send_recv_cu(num_workers, endpoints_per_worker): +def _test_send_recv_cu( + num_workers, endpoints_per_worker, enable_monitor, communication +): ctx = multiprocessing.get_context("spawn") - enable_monitor = False monitor_port = 0 signal = ctx.Array("i", [0, 0]) ports = ctx.Array("i", range(num_workers)) lock = ctx.Lock() - comm_type = base_worker - # comm_type = asyncio_worker - # comm_type = tornado_worker - if enable_monitor: monitor_process = ctx.Process( name="worker", - target=comm_type, + target=communication, args=[signal, ports, lock, 0, num_workers, endpoints_per_worker, True, 0], ) monitor_process.start() @@ -508,7 +496,7 @@ def _test_send_recv_cu(num_workers, endpoints_per_worker): for worker_num in range(num_workers): worker_process = ctx.Process( name="worker", - target=comm_type, + target=communication, args=[ signal, ports, @@ -534,5 +522,9 @@ def _test_send_recv_cu(num_workers, endpoints_per_worker): @pytest.mark.parametrize("num_workers", [2, 4, 8]) @pytest.mark.parametrize("endpoints_per_worker", [1]) -def test_send_recv_cu(num_workers, endpoints_per_worker): - _test_send_recv_cu(num_workers, endpoints_per_worker) +@pytest.mark.parametrize("enable_monitor", [True, False]) +@pytest.mark.parametrize( + "communication", [ucx_process, asyncio_process, tornado_process] +) +def test_send_recv_cu(num_workers, endpoints_per_worker, enable_monitor, communication): + _test_send_recv_cu(num_workers, endpoints_per_worker, enable_monitor, communication) diff --git a/tests/utils.py b/tests/utils.py index bbf0aa409..36e8fefde 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -235,7 +235,7 @@ async def recv(self): ) except EOFError: raise Exception("aborted stream on truncated data") - return msg + return frames, msg async def close(self): self.stream.close() @@ -261,7 +261,7 @@ async def start_server(cls, host, port): server = TCPServer(max_buffer_size=2 ** 30) - if port is None: + if port is None or port == 0: def _try_listen(server, host): while True: @@ -313,10 +313,8 @@ async def send(self, message): nframes = len(frames) self.writer.write(struct.pack("Q", nframes)) - # await self.writer.drain() sizes = list(nbytes(f) for f in frames) self.writer.write(struct.pack(nframes * "Q", *sizes)) - # await self.writer.drain() for f in frames: self.writer.write(f) await self.writer.drain() @@ -390,12 +388,10 @@ async def send(self, message): frames = await to_frames(msg, serializers=("cuda", "dask", "pickle")) nframes = len(frames) - print(f"Send nframes: {nframes}") 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)) - print(f"Send sizes: {sizes}") for f in frames: await self.ep.send(f) @@ -404,12 +400,10 @@ async def recv(self): nframes = np.empty((struct.calcsize("Q"),), dtype="u1") await self.ep.recv(nframes) nframes = struct.unpack("Q", nframes) - print(f"Recv nframes: {nframes}") sizes = np.empty((struct.calcsize(nframes[0] * "Q"),), dtype="u1") await self.ep.recv(sizes) sizes = struct.unpack(nframes[0] * "Q", sizes) - print(f"Recv sizes: {sizes}") frames = [] for size in sizes: From e21a9732b4821df1f3110ff1c28a33f2e85d98dc Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Thu, 8 Jul 2021 14:30:28 -0700 Subject: [PATCH 21/42] Remove worker_num argument --- tests/test_multiple_processes_all_to_all.py | 74 +++------------------ 1 file changed, 10 insertions(+), 64 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index f84e4f6b3..b40de325e 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -37,7 +37,6 @@ def __init__( signal, ports, lock, - worker_num, num_workers, endpoints_per_worker, monitor_port, @@ -46,7 +45,6 @@ def __init__( self.signal = signal self.ports = ports self.lock = lock - self.worker_num = worker_num self.num_workers = num_workers self.endpoints_per_worker = endpoints_per_worker @@ -184,8 +182,8 @@ async def run_monitor(self): async def _wait_for_workers(self): if self.monitor_port == 0: with self.lock: + self.ports[self.signal[0]] = self.listener.port self.signal[0] += 1 - self.ports[self.worker_num] = self.listener.port while self.signal[0] != self.num_workers: await self._sleep(0.1) @@ -204,8 +202,8 @@ async def _wait_for_workers(self): async def _wait_for_completion(self): if self.monitor_port == 0: with self.lock: + self.ports[self.signal[1]] = self.listener.port self.signal[1] += 1 - self.ports[self.worker_num] = self.listener.port while self.signal[1] != self.num_workers: await self._sleep(0) @@ -326,20 +324,12 @@ def get_results(self): class TornadoProcess(UCXProcess): def __init__( - self, - signal, - ports, - lock, - worker_num, - num_workers, - endpoints_per_worker, - monitor_port, + self, signal, ports, lock, num_workers, endpoints_per_worker, monitor_port, ): super().__init__( signal, ports, lock, - worker_num, num_workers, endpoints_per_worker, monitor_port, @@ -364,20 +354,12 @@ async def _create_endpoint(self, host, port): class AsyncioProcess(UCXProcess): def __init__( - self, - signal, - ports, - lock, - worker_num, - num_workers, - endpoints_per_worker, - monitor_port, + self, signal, ports, lock, num_workers, endpoints_per_worker, monitor_port, ): super().__init__( signal, ports, lock, - worker_num, num_workers, endpoints_per_worker, monitor_port, @@ -391,25 +373,16 @@ async def _create_listener(self, host, port=None, cb=None): return await AsyncioCommServer.start_server(host, port) async def _create_endpoint(self, host, port): - host = ucp.get_address(ifname="enp1s0f0") return await AsyncioCommConnection.open_connection(host, port) def ucx_process( - signal, - ports, - lock, - worker_num, - num_workers, - endpoints_per_worker, - is_monitor, - monitor_port, + signal, ports, lock, num_workers, endpoints_per_worker, is_monitor, monitor_port, ): w = UCXProcess( signal, ports, lock, - worker_num, num_workers, endpoints_per_worker, monitor_port, @@ -421,23 +394,10 @@ def ucx_process( def asyncio_process( - signal, - ports, - lock, - worker_num, - num_workers, - endpoints_per_worker, - is_monitor, - monitor_port, + signal, ports, lock, num_workers, endpoints_per_worker, is_monitor, monitor_port, ): w = AsyncioProcess( - signal, - ports, - lock, - worker_num, - num_workers, - endpoints_per_worker, - monitor_port, + signal, ports, lock, num_workers, endpoints_per_worker, monitor_port, ) run_func = w.run_monitor if is_monitor else w.run asyncio.get_event_loop().run_until_complete(run_func()) @@ -445,23 +405,10 @@ def asyncio_process( def tornado_process( - signal, - ports, - lock, - worker_num, - num_workers, - endpoints_per_worker, - is_monitor, - monitor_port, + signal, ports, lock, num_workers, endpoints_per_worker, is_monitor, monitor_port, ): w = TornadoProcess( - signal, - ports, - lock, - worker_num, - num_workers, - endpoints_per_worker, - monitor_port, + signal, ports, lock, num_workers, endpoints_per_worker, monitor_port, ) run_func = w.run_monitor if is_monitor else w.run IOLoop.current().run_sync(run_func) @@ -483,7 +430,7 @@ def _test_send_recv_cu( monitor_process = ctx.Process( name="worker", target=communication, - args=[signal, ports, lock, 0, num_workers, endpoints_per_worker, True, 0], + args=[signal, ports, lock, num_workers, endpoints_per_worker, True, 0], ) monitor_process.start() @@ -501,7 +448,6 @@ def _test_send_recv_cu( signal, ports, lock, - worker_num, num_workers, endpoints_per_worker, False, From 6c7de452999a5689fd31b6999e8d94f06dc3434b Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Thu, 8 Jul 2021 15:39:00 -0700 Subject: [PATCH 22/42] Make shm synchronization arguments optional --- tests/test_multiple_processes_all_to_all.py | 127 ++++++++++++++------ 1 file changed, 91 insertions(+), 36 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index b40de325e..f46a4346a 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -34,17 +34,15 @@ class UCXProcess: def __init__( self, - signal, - ports, - lock, num_workers, endpoints_per_worker, monitor_port, transfer_to_cache, + shm_sync=True, + signal=None, + ports=None, + lock=None, ): - self.signal = signal - self.ports = ports - self.lock = lock self.num_workers = num_workers self.endpoints_per_worker = endpoints_per_worker @@ -54,6 +52,11 @@ def __init__( self.monitor_port = monitor_port + self.shm_sync = shm_sync + self.signal = signal + self.ports = ports + self.lock = lock + self.conns = dict() self.connections = set() self.bytes_bandwidth = dict() @@ -135,8 +138,11 @@ async def run_monitor(self): self.listener_address, self.monitor_port, self._monitor_listener_cb, ) - with self.lock: - self.signal[0] = self.listener.port + 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()) != ( @@ -180,7 +186,7 @@ async def run_monitor(self): pass async def _wait_for_workers(self): - if self.monitor_port == 0: + if self.shm_sync: with self.lock: self.ports[self.signal[0]] = self.listener.port self.signal[0] += 1 @@ -200,7 +206,7 @@ async def _wait_for_workers(self): assert len(self.worker_addresses) == self.num_workers async def _wait_for_completion(self): - if self.monitor_port == 0: + if self.shm_sync: with self.lock: self.ports[self.signal[1]] = self.listener.port self.signal[1] += 1 @@ -324,16 +330,24 @@ def get_results(self): class TornadoProcess(UCXProcess): def __init__( - self, signal, ports, lock, num_workers, endpoints_per_worker, monitor_port, + self, + num_workers, + endpoints_per_worker, + monitor_port, + shm_sync=True, + signal=None, + ports=None, + lock=None, ): super().__init__( - signal, - ports, - lock, num_workers, endpoints_per_worker, monitor_port, transfer_to_cache=False, + shm_sync=shm_sync, + signal=signal, + ports=ports, + lock=lock, ) async def _sleep(self, delay): @@ -354,16 +368,24 @@ async def _create_endpoint(self, host, port): class AsyncioProcess(UCXProcess): def __init__( - self, signal, ports, lock, num_workers, endpoints_per_worker, monitor_port, + self, + num_workers, + endpoints_per_worker, + monitor_port, + shm_sync=True, + signal=None, + ports=None, + lock=None, ): super().__init__( - signal, - ports, - lock, num_workers, endpoints_per_worker, monitor_port, transfer_to_cache=False, + shm_sync=shm_sync, + signal=signal, + ports=ports, + lock=lock, ) def _init(self): @@ -377,16 +399,24 @@ async def _create_endpoint(self, host, port): def ucx_process( - signal, ports, lock, num_workers, endpoints_per_worker, is_monitor, monitor_port, + num_workers, + endpoints_per_worker, + is_monitor, + monitor_port, + shm_sync=True, + signal=None, + ports=None, + lock=None, ): w = UCXProcess( - signal, - ports, - lock, num_workers, endpoints_per_worker, monitor_port, 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 asyncio.get_event_loop().run_until_complete(run_func()) @@ -394,10 +424,23 @@ def ucx_process( def asyncio_process( - signal, ports, lock, num_workers, endpoints_per_worker, is_monitor, monitor_port, + num_workers, + endpoints_per_worker, + is_monitor, + monitor_port, + shm_sync=True, + signal=None, + ports=None, + lock=None, ): w = AsyncioProcess( - signal, ports, lock, num_workers, endpoints_per_worker, monitor_port, + num_workers, + endpoints_per_worker, + monitor_port, + shm_sync=shm_sync, + signal=signal, + ports=ports, + lock=lock, ) run_func = w.run_monitor if is_monitor else w.run asyncio.get_event_loop().run_until_complete(run_func()) @@ -405,10 +448,23 @@ def asyncio_process( def tornado_process( - signal, ports, lock, num_workers, endpoints_per_worker, is_monitor, monitor_port, + num_workers, + endpoints_per_worker, + is_monitor, + monitor_port, + shm_sync=True, + signal=None, + ports=None, + lock=None, ): w = TornadoProcess( - signal, ports, lock, num_workers, endpoints_per_worker, monitor_port, + num_workers, + endpoints_per_worker, + monitor_port, + shm_sync=shm_sync, + signal=signal, + ports=ports, + lock=lock, ) run_func = w.run_monitor if is_monitor else w.run IOLoop.current().run_sync(run_func) @@ -430,7 +486,8 @@ def _test_send_recv_cu( monitor_process = ctx.Process( name="worker", target=communication, - args=[signal, ports, lock, num_workers, endpoints_per_worker, True, 0], + args=[num_workers, endpoints_per_worker, True, 0], + kwargs={"shm_sync": True, "signal": signal, "ports": ports, "lock": lock}, ) monitor_process.start() @@ -444,15 +501,13 @@ def _test_send_recv_cu( worker_process = ctx.Process( name="worker", target=communication, - args=[ - signal, - ports, - lock, - num_workers, - endpoints_per_worker, - False, - monitor_port, - ], + args=[num_workers, endpoints_per_worker, False, monitor_port], + kwargs={ + "shm_sync": not enable_monitor, + "signal": signal, + "ports": ports, + "lock": lock, + }, ) worker_process.start() worker_processes.append(worker_process) From 1bdc4d23c1099bf42cf5e72c8d3f3089faf3e098 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Thu, 8 Jul 2021 15:52:52 -0700 Subject: [PATCH 23/42] Add listener_address argument --- tests/test_multiple_processes_all_to_all.py | 25 +++++++++++++++++---- 1 file changed, 21 insertions(+), 4 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index f46a4346a..006e8dfcb 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -34,6 +34,7 @@ class UCXProcess: def __init__( self, + listener_address, num_workers, endpoints_per_worker, monitor_port, @@ -43,6 +44,7 @@ def __init__( ports=None, lock=None, ): + self.listener_address = listener_address self.num_workers = num_workers self.endpoints_per_worker = endpoints_per_worker @@ -133,7 +135,6 @@ def get_connections(self): async def run_monitor(self): self._init() - self.listener_address = ucp.get_address(ifname="enp1s0f0") self.listener = await self._create_listener( self.listener_address, self.monitor_port, self._monitor_listener_cb, ) @@ -295,7 +296,6 @@ async def run(self): self._init() # Start listener - self.listener_address = ucp.get_address(ifname="enp1s0f0") self.listener = await self._create_listener(self.listener_address) await self._wait_for_workers() @@ -331,6 +331,7 @@ def get_results(self): class TornadoProcess(UCXProcess): def __init__( self, + listener_address, num_workers, endpoints_per_worker, monitor_port, @@ -340,6 +341,7 @@ def __init__( lock=None, ): super().__init__( + listener_address, num_workers, endpoints_per_worker, monitor_port, @@ -369,6 +371,7 @@ async def _create_endpoint(self, host, port): class AsyncioProcess(UCXProcess): def __init__( self, + listener_address, num_workers, endpoints_per_worker, monitor_port, @@ -378,6 +381,7 @@ def __init__( lock=None, ): super().__init__( + listener_address, num_workers, endpoints_per_worker, monitor_port, @@ -399,6 +403,7 @@ async def _create_endpoint(self, host, port): def ucx_process( + listener_address, num_workers, endpoints_per_worker, is_monitor, @@ -409,6 +414,7 @@ def ucx_process( lock=None, ): w = UCXProcess( + listener_address, num_workers, endpoints_per_worker, monitor_port, @@ -424,6 +430,7 @@ def ucx_process( def asyncio_process( + listener_address, num_workers, endpoints_per_worker, is_monitor, @@ -434,6 +441,7 @@ def asyncio_process( lock=None, ): w = AsyncioProcess( + listener_address, num_workers, endpoints_per_worker, monitor_port, @@ -448,6 +456,7 @@ def asyncio_process( def tornado_process( + listener_address, num_workers, endpoints_per_worker, is_monitor, @@ -458,6 +467,7 @@ def tornado_process( lock=None, ): w = TornadoProcess( + listener_address, num_workers, endpoints_per_worker, monitor_port, @@ -476,6 +486,7 @@ def _test_send_recv_cu( ): ctx = multiprocessing.get_context("spawn") + listener_address = ucp.get_address(ifname="enp1s0f0") monitor_port = 0 signal = ctx.Array("i", [0, 0]) @@ -486,7 +497,7 @@ def _test_send_recv_cu( monitor_process = ctx.Process( name="worker", target=communication, - args=[num_workers, endpoints_per_worker, True, 0], + args=[listener_address, num_workers, endpoints_per_worker, True, 0], kwargs={"shm_sync": True, "signal": signal, "ports": ports, "lock": lock}, ) monitor_process.start() @@ -501,7 +512,13 @@ def _test_send_recv_cu( worker_process = ctx.Process( name="worker", target=communication, - args=[num_workers, endpoints_per_worker, False, monitor_port], + args=[ + listener_address, + num_workers, + endpoints_per_worker, + False, + monitor_port, + ], kwargs={ "shm_sync": not enable_monitor, "signal": signal, From da75bc1641c120bb0959a4a6c5821e17a28dab14 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Fri, 9 Jul 2021 07:29:42 -0700 Subject: [PATCH 24/42] Separate benchmark into multiple files --- ...benchmark_multiple_processes_all_to_all.py | 163 ++++++ tests/test_multiple_processes_all_to_all.py | 477 +---------------- tests/utils.py | 320 +----------- tests/utils_all_to_all.py | 483 ++++++++++++++++++ tests/utils_comm_libs.py | 325 ++++++++++++ 5 files changed, 973 insertions(+), 795 deletions(-) create mode 100644 tests/benchmark_multiple_processes_all_to_all.py create mode 100644 tests/utils_all_to_all.py create mode 100644 tests/utils_comm_libs.py 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..78ff5e08c --- /dev/null +++ b/tests/benchmark_multiple_processes_all_to_all.py @@ -0,0 +1,163 @@ +import argparse +import multiprocessing + +from utils_all_to_all import asyncio_process, tornado_process, ucx_process + +import ucp + + +def parse_args(): + parser = argparse.ArgumentParser(description="All-to-all benchmark") + parser.add_argument( + "--multi-node", default=False, action="store_true", + ) + parser.add_argument( + "--monitor", default=False, action="store_true", + ) + parser.add_argument( + "--worker", default=False, action="store_true", + ) + parser.add_argument( + "--enable-monitor", + default=False, + action="store_true", + help="Use a monitor process to synchronize workers (always " + "enabled with --multi-node), otherwise synchronize processes " + "via multiprocessing shared memory.", + ) + parser.add_argument( + "--listen-interface", + default=None, + help="Interface where monitor (if --monitor), or worker " + "(if --worker), or all (if not --multi-node) will listen for " + "connections", + ) + parser.add_argument( + "--monitor-address", + default=None, + help="Address where monitor is listening to process is started " + "with --worker, in the HOST:PORT format", + ) + parser.add_argument( + "--num-workers", default=2, type=int, + ) + parser.add_argument( + "--endpoints-per-worker", default=1, type=int, + ) + parser.add_argument( + "--communication-lib", + default="ucx", + type=str, + help="Communication library to benchmark. Options are " + "'ucx' (default), 'asyncio', 'tornado'.", + ) + + args = parser.parse_args() + if args.worker and args.monitor: + raise RuntimeError("--monitor and --worker can't be defined together.") + if args.multi_node and not (args.worker or args.monitor): + raise RuntimeError( + "Either --monitor or --worker need to be defined together with " + "--multi-node." + ) + 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 == "tornado": + communication_func = tornado_process + else: + raise ValueError( + f"Communication library {args.communication_lib} not supported" + ) + + if args.multi_node is False: + ctx = multiprocessing.get_context("spawn") + + signal = ctx.Array("i", [0, 0]) + ports = ctx.Array("i", range(num_workers)) + lock = ctx.Lock() + + if args.enable_monitor: + monitor_process = ctx.Process( + name="worker", + target=communication_func, + args=[listener_address, num_workers, endpoints_per_worker, True, 0], + 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, + ], + 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() + + monitor_process.join() + + assert worker_process.exitcode == 0 + else: + if args.monitor is True: + communication_func( + listener_address, + num_workers, + endpoints_per_worker, + True, + 0, + shm_sync=False, + ) + elif args.worker is True: + communication_func( + listener_address, + num_workers, + endpoints_per_worker, + False, + monitor_port, + shm_sync=False, + ) + + +if __name__ == "__main__": + main() diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index 006e8dfcb..cfe5ccd7c 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -1,485 +1,10 @@ -import asyncio import multiprocessing -from time import monotonic -import numpy as np import pytest -from tornado import gen -from tornado.ioloop import IOLoop -from utils import ( - AsyncioCommConnection, - AsyncioCommServer, - TornadoTCPConnection, - TornadoTCPServer, - UCXConnection, - UCXServer, -) - -from dask.utils import format_bytes -from distributed.utils import nbytes +from utils_all_to_all import asyncio_process, tornado_process, ucx_process import ucp -GatherSendRecv = False -Iterations = 3 -Size = 2 ** 20 - -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, - 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 - - # UCX only effectively creates endpoints at first transfer, but this - # isn't necessary for tornado/asyncio. - self.transfer_to_cache = transfer_to_cache - - self.monitor_port = monitor_port - - 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): - for i in range(Iterations): - if GatherSendRecv: - 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) - - async def _listener(self, ep): - message = np.arange(Size, dtype=np.uint8) - - await self._transfer(ep, message) - - async def _client(self, my_port, worker_address, ep, cache_only=False): - message = np.arange(Size, dtype=np.uint8) - send_recv_bytes = (nbytes(message) * 2) * Iterations - - t = monotonic() - await self._transfer(ep, message, send_first=False) - total_time = monotonic() - t - - if cache_only is False: - self.bytes_bandwidth[worker_address].append( - (send_recv_bytes, send_recv_bytes / total_time) - ) - - 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): - pass - - 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 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.endpoints_per_worker * self.num_workers - ): - await self._sleep(0.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,)) - - 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 _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] - assert len(self.worker_addresses) == self.num_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.monitor_port == 0: - for remote_port in list(self.ports): - if remote_port == self.listener.port: - continue - - ep = await self._create_endpoint(self.listener_address, remote_port) - self.bytes_bandwidth[remote_port] = [] - self.conns[(remote_port, i)] = ep - - if self.transfer_to_cache: - client_tasks.append( - self._client( - self.listener.port, remote_port, 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(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() - - def get_results(self): - for remote_port, bb in self.bytes_bandwidth.items(): - 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( - "[%d, %s] Transferred bytes: %s, average bandwidth: %s/s, " - "median bandwidth: %s/s" - % ( - self.listener.port, - remote_port, - 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, - shm_sync=True, - signal=None, - ports=None, - lock=None, - ): - super().__init__( - listener_address, - num_workers, - endpoints_per_worker, - monitor_port, - 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, - shm_sync=True, - signal=None, - ports=None, - lock=None, - ): - super().__init__( - listener_address, - num_workers, - endpoints_per_worker, - monitor_port, - 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, - shm_sync=True, - signal=None, - ports=None, - lock=None, -): - w = UCXProcess( - listener_address, - num_workers, - endpoints_per_worker, - monitor_port, - 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 - 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, - shm_sync=True, - signal=None, - ports=None, - lock=None, -): - w = AsyncioProcess( - listener_address, - num_workers, - endpoints_per_worker, - monitor_port, - shm_sync=shm_sync, - signal=signal, - ports=ports, - lock=lock, - ) - run_func = w.run_monitor if is_monitor else w.run - asyncio.get_event_loop().run_until_complete(run_func()) - w.get_results() - - -def tornado_process( - listener_address, - num_workers, - endpoints_per_worker, - is_monitor, - monitor_port, - shm_sync=True, - signal=None, - ports=None, - lock=None, -): - w = TornadoProcess( - listener_address, - num_workers, - endpoints_per_worker, - monitor_port, - shm_sync=shm_sync, - signal=signal, - ports=ports, - lock=lock, - ) - run_func = w.run_monitor if is_monitor else w.run - IOLoop.current().run_sync(run_func) - w.get_results() - def _test_send_recv_cu( num_workers, endpoints_per_worker, enable_monitor, communication diff --git a/tests/utils.py b/tests/utils.py index 36e8fefde..ef8ecccd7 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -1,19 +1,11 @@ -import asyncio -import functools import io import logging import os -import struct from contextlib import contextmanager 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.comm.utils import from_frames from distributed.utils import nbytes import rmm @@ -146,313 +138,3 @@ async def am_recv(ep): msg = await from_frames(frames) return frames, msg - - -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): - def _server_callback(connections, reader, writer): - connections.append(AsyncioCommConnection(reader, writer)) - - connections = [] - - server = await asyncio.start_server( - functools.partial(_server_callback, connections), 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): - async def _server_callback(connections, ep): - conn = UCXConnection(ep) - connections.append(UCXConnection(ep)) - await listener_func(conn) - - connections = [] - - server = ucp.create_listener( - functools.partial(_server_callback, connections), 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() diff --git a/tests/utils_all_to_all.py b/tests/utils_all_to_all.py new file mode 100644 index 000000000..5fd214b47 --- /dev/null +++ b/tests/utils_all_to_all.py @@ -0,0 +1,483 @@ +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 + +GatherSendRecv = False +Iterations = 3 +Size = 2 ** 20 + +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, + 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 + + # UCX only effectively creates endpoints at first transfer, but this + # isn't necessary for tornado/asyncio. + self.transfer_to_cache = transfer_to_cache + + self.monitor_port = monitor_port + + 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): + for i in range(Iterations): + if GatherSendRecv: + 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) + + async def _listener(self, ep): + message = np.arange(Size, dtype=np.uint8) + + await self._transfer(ep, message) + + async def _client(self, my_port, worker_address, ep, cache_only=False): + message = np.arange(Size, dtype=np.uint8) + send_recv_bytes = (nbytes(message) * 2) * Iterations + + t = monotonic() + await self._transfer(ep, message, send_first=False) + total_time = monotonic() - t + + if cache_only is False: + self.bytes_bandwidth[worker_address].append( + (send_recv_bytes, send_recv_bytes / total_time) + ) + + 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): + pass + + 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 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: + print( + f"Waiting for all workers to connect, {len(self.get_connections())} " + f"of {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,)) + + 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 _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] + assert len(self.worker_addresses) == self.num_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.monitor_port == 0: + for remote_port in list(self.ports): + if remote_port == self.listener.port: + continue + + ep = await self._create_endpoint(self.listener_address, remote_port) + self.bytes_bandwidth[remote_port] = [] + self.conns[(remote_port, i)] = ep + + if self.transfer_to_cache: + client_tasks.append( + self._client( + self.listener.port, remote_port, 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(self): + self._init() + + # Start listener + self.listener = await self._create_listener(self.listener_address) + + print("Wait for workers") + await self._wait_for_workers() + + print("Create endpoints") + 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() + + def get_results(self): + for remote_port, bb in self.bytes_bandwidth.items(): + 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( + "[%d, %s] Transferred bytes: %s, average bandwidth: %s/s, " + "median bandwidth: %s/s" + % ( + self.listener.port, + remote_port, + 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, + shm_sync=True, + signal=None, + ports=None, + lock=None, + ): + super().__init__( + listener_address, + num_workers, + endpoints_per_worker, + monitor_port, + 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, + shm_sync=True, + signal=None, + ports=None, + lock=None, + ): + super().__init__( + listener_address, + num_workers, + endpoints_per_worker, + monitor_port, + 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, + shm_sync=True, + signal=None, + ports=None, + lock=None, +): + w = UCXProcess( + listener_address, + num_workers, + endpoints_per_worker, + monitor_port, + 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 + 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, + shm_sync=True, + signal=None, + ports=None, + lock=None, +): + w = AsyncioProcess( + listener_address, + num_workers, + endpoints_per_worker, + monitor_port, + shm_sync=shm_sync, + signal=signal, + ports=ports, + lock=lock, + ) + run_func = w.run_monitor if is_monitor else w.run + asyncio.get_event_loop().run_until_complete(run_func()) + w.get_results() + + +def tornado_process( + listener_address, + num_workers, + endpoints_per_worker, + is_monitor, + monitor_port, + shm_sync=True, + signal=None, + ports=None, + lock=None, +): + w = TornadoProcess( + listener_address, + num_workers, + endpoints_per_worker, + monitor_port, + shm_sync=shm_sync, + signal=signal, + ports=ports, + lock=lock, + ) + run_func = w.run_monitor if is_monitor else w.run + 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..c7bdafa9f --- /dev/null +++ b/tests/utils_comm_libs.py @@ -0,0 +1,325 @@ +import asyncio +import functools +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): + def _server_callback(connections, reader, writer): + connections.append(AsyncioCommConnection(reader, writer)) + + connections = [] + + server = await asyncio.start_server( + functools.partial(_server_callback, connections), 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): + async def _server_callback(connections, ep): + conn = UCXConnection(ep) + connections.append(conn) + await listener_func(conn) + + connections = [] + + server = ucp.create_listener( + functools.partial(_server_callback, connections), 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() From 92023100a072a69c3770988c11526145436bc666 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Fri, 9 Jul 2021 07:41:10 -0700 Subject: [PATCH 25/42] Improve benchmark result formatting --- tests/utils_all_to_all.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/tests/utils_all_to_all.py b/tests/utils_all_to_all.py index 5fd214b47..acada1da0 100644 --- a/tests/utils_all_to_all.py +++ b/tests/utils_all_to_all.py @@ -298,10 +298,8 @@ async def run(self): # Start listener self.listener = await self._create_listener(self.listener_address) - print("Wait for workers") await self._wait_for_workers() - print("Create endpoints") await self._create_all_endpoints() await self._wait_for_connections_cache() @@ -314,14 +312,15 @@ async def run(self): def get_results(self): for remote_port, 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( - "[%d, %s] Transferred bytes: %s, average bandwidth: %s/s, " + "[%s -> %s] Transferred bytes: %s, average bandwidth: %s/s, " "median bandwidth: %s/s" % ( - self.listener.port, + local_address, remote_port, format_bytes(total_bytes), format_bytes(avg_bandwidth), From dba6dbfa339685cba332d8d263be4c506df2e479 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Fri, 9 Jul 2021 10:26:25 -0700 Subject: [PATCH 26/42] Reorganize run_worker and run_monitor --- tests/utils_all_to_all.py | 112 +++++++++++++++++--------------------- 1 file changed, 51 insertions(+), 61 deletions(-) diff --git a/tests/utils_all_to_all.py b/tests/utils_all_to_all.py index acada1da0..c1b3a563f 100644 --- a/tests/utils_all_to_all.py +++ b/tests/utils_all_to_all.py @@ -108,7 +108,7 @@ 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): - pass + return None async def _create_endpoint(self, host, port): my_port = self.listener.port @@ -130,62 +130,6 @@ async def _create_endpoint(self, host, port): def get_connections(self): return self.listener._connections - 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: - print( - f"Waiting for all workers to connect, {len(self.get_connections())} " - f"of {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,)) - - 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 _wait_for_workers(self): if self.shm_sync: with self.lock: @@ -292,7 +236,7 @@ async def _close_connections_and_listener(self): except ucp.UCXCloseError: pass - async def run(self): + async def run_worker(self): self._init() # Start listener @@ -310,6 +254,52 @@ async def run(self): 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: + print( + f"Waiting for all workers to connect, {len(self.get_connections())} " + f"of {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_port, bb in self.bytes_bandwidth.items(): local_address = (self.listener_address, self.listener.port) @@ -425,7 +415,7 @@ def ucx_process( ports=ports, lock=lock, ) - run_func = w.run_monitor if is_monitor else w.run + run_func = w.run_monitor if is_monitor else w.run_worker asyncio.get_event_loop().run_until_complete(run_func()) w.get_results() @@ -451,7 +441,7 @@ def asyncio_process( ports=ports, lock=lock, ) - run_func = w.run_monitor if is_monitor else w.run + run_func = w.run_monitor if is_monitor else w.run_worker asyncio.get_event_loop().run_until_complete(run_func()) w.get_results() @@ -477,6 +467,6 @@ def tornado_process( ports=ports, lock=lock, ) - run_func = w.run_monitor if is_monitor else w.run + run_func = w.run_monitor if is_monitor else w.run_worker IOLoop.current().run_sync(run_func) w.get_results() From 26c371297aba1d3773173e0f5b74055031b95e8b Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Fri, 9 Jul 2021 11:50:05 -0700 Subject: [PATCH 27/42] Fix single-node benchmark --- tests/benchmark_multiple_processes_all_to_all.py | 5 ++++- tests/utils_all_to_all.py | 2 +- 2 files changed, 5 insertions(+), 2 deletions(-) diff --git a/tests/benchmark_multiple_processes_all_to_all.py b/tests/benchmark_multiple_processes_all_to_all.py index 78ff5e08c..0b7bb2fa2 100644 --- a/tests/benchmark_multiple_processes_all_to_all.py +++ b/tests/benchmark_multiple_processes_all_to_all.py @@ -91,6 +91,8 @@ def main(): ports = ctx.Array("i", range(num_workers)) lock = ctx.Lock() + monitor_port = 0 + if args.enable_monitor: monitor_process = ctx.Process( name="worker", @@ -135,7 +137,8 @@ def main(): for worker_process in worker_processes: worker_process.join() - monitor_process.join() + if args.enable_monitor: + monitor_process.join() assert worker_process.exitcode == 0 else: diff --git a/tests/utils_all_to_all.py b/tests/utils_all_to_all.py index c1b3a563f..e365b2d18 100644 --- a/tests/utils_all_to_all.py +++ b/tests/utils_all_to_all.py @@ -168,7 +168,7 @@ async def _create_all_endpoints(self): for i in range(self.endpoints_per_worker): client_tasks = [] # Create endpoints to all other workers - if self.monitor_port == 0: + if self.shm_sync: for remote_port in list(self.ports): if remote_port == self.listener.port: continue From 199f960fab8a1cbd19d913582a4deca12b638857 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Fri, 9 Jul 2021 12:07:52 -0700 Subject: [PATCH 28/42] Add more all-to-all benchmark CLI arguments --- ...benchmark_multiple_processes_all_to_all.py | 42 +++++++++++- tests/utils_all_to_all.py | 68 ++++++++++++++----- 2 files changed, 91 insertions(+), 19 deletions(-) diff --git a/tests/benchmark_multiple_processes_all_to_all.py b/tests/benchmark_multiple_processes_all_to_all.py index 0b7bb2fa2..3ef4b76c8 100644 --- a/tests/benchmark_multiple_processes_all_to_all.py +++ b/tests/benchmark_multiple_processes_all_to_all.py @@ -51,6 +51,28 @@ def parse_args(): help="Communication library to benchmark. Options are " "'ucx' (default), 'asyncio', 'tornado'.", ) + parser.add_argument( + "--size", + default=2 ** 20, + type=int, + help="Size to be passed for data generation function", + ) + 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.", + ) + 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.", + ) args = parser.parse_args() if args.worker and args.monitor: @@ -97,7 +119,16 @@ def main(): monitor_process = ctx.Process( name="worker", target=communication_func, - args=[listener_address, num_workers, endpoints_per_worker, True, 0], + args=[ + listener_address, + num_workers, + endpoints_per_worker, + True, + 0, + args.size, + args.iterations, + args.gather_send_recv, + ], kwargs={ "shm_sync": True, "signal": signal, @@ -123,6 +154,9 @@ def main(): endpoints_per_worker, False, monitor_port, + args.size, + args.iterations, + args.gather_send_recv, ], kwargs={ "shm_sync": not args.enable_monitor, @@ -149,6 +183,9 @@ def main(): endpoints_per_worker, True, 0, + args.size, + args.iterations, + args.gather_send_recv, shm_sync=False, ) elif args.worker is True: @@ -158,6 +195,9 @@ def main(): endpoints_per_worker, False, monitor_port, + args.size, + args.iterations, + args.gather_send_recv, shm_sync=False, ) diff --git a/tests/utils_all_to_all.py b/tests/utils_all_to_all.py index e365b2d18..a48c82f6d 100644 --- a/tests/utils_all_to_all.py +++ b/tests/utils_all_to_all.py @@ -18,10 +18,6 @@ import ucp -GatherSendRecv = False -Iterations = 3 -Size = 2 ** 20 - OP_NONE = 0 OP_WORKER_LISTENING = 1 OP_CLUSTER_READY = 2 @@ -36,6 +32,9 @@ def __init__( num_workers, endpoints_per_worker, monitor_port, + size, + iterations, + gather_send_recv, transfer_to_cache, shm_sync=True, signal=None, @@ -45,12 +44,15 @@ def __init__( 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.monitor_port = monitor_port + self.size = size + self.iterations = iterations + self.gather_send_recv = gather_send_recv self.shm_sync = shm_sync self.signal = signal @@ -70,8 +72,8 @@ async def _gather(self, tasks): await asyncio.gather(*tasks) async def _transfer(self, ep, msg2send, send_first=True): - for i in range(Iterations): - if GatherSendRecv: + for i in range(self.iterations): + if self.gather_send_recv: msgs = [ep.recv(), ep.send(msg2send)] await self._gather(msgs) else: @@ -84,13 +86,13 @@ async def _transfer(self, ep, msg2send, send_first=True): await ep.send(msg2send) async def _listener(self, ep): - message = np.arange(Size, dtype=np.uint8) + 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(Size, dtype=np.uint8) - send_recv_bytes = (nbytes(message) * 2) * Iterations + message = np.arange(self.size, dtype=np.uint8) + send_recv_bytes = (nbytes(message) * 2) * self.iterations t = monotonic() await self._transfer(ep, message, send_first=False) @@ -326,16 +328,22 @@ def __init__( 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, - num_workers, - endpoints_per_worker, - monitor_port, + 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, @@ -366,16 +374,22 @@ def __init__( 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, - num_workers, - endpoints_per_worker, - monitor_port, + 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, @@ -399,6 +413,9 @@ def ucx_process( endpoints_per_worker, is_monitor, monitor_port, + size, + iterations, + gather_send_recv, shm_sync=True, signal=None, ports=None, @@ -409,6 +426,9 @@ def ucx_process( 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, @@ -426,6 +446,9 @@ def asyncio_process( endpoints_per_worker, is_monitor, monitor_port, + size, + iterations, + gather_send_recv, shm_sync=True, signal=None, ports=None, @@ -436,6 +459,9 @@ def asyncio_process( 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, @@ -452,6 +478,9 @@ def tornado_process( endpoints_per_worker, is_monitor, monitor_port, + size, + iterations, + gather_send_recv, shm_sync=True, signal=None, ports=None, @@ -462,6 +491,9 @@ def tornado_process( 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, From 49b194fe053dc0de563872216b8291aec56897ae Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Fri, 9 Jul 2021 12:17:25 -0700 Subject: [PATCH 29/42] Remove hardcoded enp1s0f0 interface --- tests/test_multiple_processes_all_to_all.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index cfe5ccd7c..0af3a68d4 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -11,7 +11,7 @@ def _test_send_recv_cu( ): ctx = multiprocessing.get_context("spawn") - listener_address = ucp.get_address(ifname="enp1s0f0") + listener_address = ucp.get_address() monitor_port = 0 signal = ctx.Array("i", [0, 0]) From 796880f4fb711b08c57954226aea73c1c9c3c4d9 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Fri, 9 Jul 2021 12:32:02 -0700 Subject: [PATCH 30/42] Remove redundant --multi-node argument --- ...benchmark_multiple_processes_all_to_all.py | 59 ++++++++----------- 1 file changed, 25 insertions(+), 34 deletions(-) diff --git a/tests/benchmark_multiple_processes_all_to_all.py b/tests/benchmark_multiple_processes_all_to_all.py index 3ef4b76c8..fa23c2910 100644 --- a/tests/benchmark_multiple_processes_all_to_all.py +++ b/tests/benchmark_multiple_processes_all_to_all.py @@ -8,9 +8,6 @@ def parse_args(): parser = argparse.ArgumentParser(description="All-to-all benchmark") - parser.add_argument( - "--multi-node", default=False, action="store_true", - ) parser.add_argument( "--monitor", default=False, action="store_true", ) @@ -77,11 +74,6 @@ def parse_args(): args = parser.parse_args() if args.worker and args.monitor: raise RuntimeError("--monitor and --worker can't be defined together.") - if args.multi_node and not (args.worker or args.monitor): - raise RuntimeError( - "Either --monitor or --worker need to be defined together with " - "--multi-node." - ) return args @@ -106,7 +98,31 @@ def main(): f"Communication library {args.communication_lib} not supported" ) - if args.multi_node is False: + if args.monitor is True: + communication_func( + listener_address, + num_workers, + endpoints_per_worker, + True, + 0, + 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]) @@ -175,31 +191,6 @@ def main(): monitor_process.join() assert worker_process.exitcode == 0 - else: - if args.monitor is True: - communication_func( - listener_address, - num_workers, - endpoints_per_worker, - True, - 0, - 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, - ) if __name__ == "__main__": From cc7961be3210ae559fc7ccf92d8f0b5806a21251 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Fri, 9 Jul 2021 12:51:09 -0700 Subject: [PATCH 31/42] Improve documentation for all-to-all benchmark CLI arguments --- ...benchmark_multiple_processes_all_to_all.py | 56 +++++++++++++------ 1 file changed, 40 insertions(+), 16 deletions(-) diff --git a/tests/benchmark_multiple_processes_all_to_all.py b/tests/benchmark_multiple_processes_all_to_all.py index fa23c2910..584e5ee2b 100644 --- a/tests/benchmark_multiple_processes_all_to_all.py +++ b/tests/benchmark_multiple_processes_all_to_all.py @@ -9,50 +9,73 @@ def parse_args(): parser = argparse.ArgumentParser(description="All-to-all benchmark") parser.add_argument( - "--monitor", default=False, action="store_true", + "--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", + "--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 (always " - "enabled with --multi-node), otherwise synchronize processes " - "via multiprocessing shared memory.", + 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 (if not --multi-node) will listen for " - "connections", + 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", default=None, - help="Address where monitor is listening to process is started " - "with --worker, in the HOST:PORT format", + help="Address where --monitor process is listening to in the HOST:PORT " + "format.", ) parser.add_argument( - "--num-workers", default=2, type=int, + "--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, + "--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' (default), 'asyncio', 'tornado'.", + "'ucx', 'asyncio', 'tornado'. Default: 'ucx'.", ) parser.add_argument( "--size", default=2 ** 20, type=int, - help="Size to be passed for data generation function", + help="Size to be passed for data generation function. Default: " "1048576.", ) parser.add_argument( "--iterations", @@ -60,7 +83,7 @@ def parse_args(): type=int, help="Number of iterations of data transfers per worker pair. " "Each iteration consists in sending and receiving the same " - "data amount.", + "data amount. Default: 15.", ) parser.add_argument( "--gather-send-recv", @@ -68,7 +91,8 @@ def parse_args(): 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.", + "simultaneously via an asyncio.gather or gen.multi operation. " + "Default: disabled.", ) args = parser.parse_args() From 18a45575b319ee05a57a73255d6b6e985672d3d4 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Fri, 9 Jul 2021 12:51:31 -0700 Subject: [PATCH 32/42] Fix all-to-all benchmark test --- tests/test_multiple_processes_all_to_all.py | 31 ++++++++++++++++++--- 1 file changed, 27 insertions(+), 4 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index 0af3a68d4..dcfdec88e 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -7,7 +7,7 @@ def _test_send_recv_cu( - num_workers, endpoints_per_worker, enable_monitor, communication + num_workers, endpoints_per_worker, enable_monitor, size, iterations, communication ): ctx = multiprocessing.get_context("spawn") @@ -22,7 +22,16 @@ def _test_send_recv_cu( monitor_process = ctx.Process( name="worker", target=communication, - args=[listener_address, num_workers, endpoints_per_worker, True, 0], + 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() @@ -43,6 +52,9 @@ def _test_send_recv_cu( endpoints_per_worker, False, monitor_port, + size, + iterations, + False, ], kwargs={ "shm_sync": not enable_monitor, @@ -66,8 +78,19 @@ def _test_send_recv_cu( @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, tornado_process] ) -def test_send_recv_cu(num_workers, endpoints_per_worker, enable_monitor, communication): - _test_send_recv_cu(num_workers, endpoints_per_worker, enable_monitor, communication) +def test_send_recv_cu( + num_workers, endpoints_per_worker, enable_monitor, size, iterations, communication +): + _test_send_recv_cu( + num_workers, + endpoints_per_worker, + enable_monitor, + size, + iterations, + communication, + ) From 58a30f84b9c967e7d9a4b3a63923b8cfa3658ead Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Fri, 9 Jul 2021 13:14:00 -0700 Subject: [PATCH 33/42] Change worker address formatting in benchmark output --- tests/utils_all_to_all.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/utils_all_to_all.py b/tests/utils_all_to_all.py index a48c82f6d..75a57ee2f 100644 --- a/tests/utils_all_to_all.py +++ b/tests/utils_all_to_all.py @@ -312,8 +312,8 @@ def get_results(self): "[%s -> %s] Transferred bytes: %s, average bandwidth: %s/s, " "median bandwidth: %s/s" % ( - local_address, - remote_port, + ":".join([str(i) for i in local_address]), + ":".join([str(i) for i in remote_port]), format_bytes(total_bytes), format_bytes(avg_bandwidth), format_bytes(median_bandwidth), From 5a509e3f0c8b0c47cd660ace3e43159f894ba5c4 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Fri, 9 Jul 2021 13:32:09 -0700 Subject: [PATCH 34/42] Add --port argument to specify monitor port --- tests/benchmark_multiple_processes_all_to_all.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/tests/benchmark_multiple_processes_all_to_all.py b/tests/benchmark_multiple_processes_all_to_all.py index 584e5ee2b..77c3ec4c8 100644 --- a/tests/benchmark_multiple_processes_all_to_all.py +++ b/tests/benchmark_multiple_processes_all_to_all.py @@ -45,6 +45,14 @@ def parse_args(): 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, @@ -128,7 +136,7 @@ def main(): num_workers, endpoints_per_worker, True, - 0, + args.port, args.size, args.iterations, args.gather_send_recv, From 75fca3d8e3f0ebb7ff0834b29c1509a3b6c4653e Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Fri, 9 Jul 2021 14:29:46 -0700 Subject: [PATCH 35/42] Fix issues with all-to-all pytest --- tests/utils_all_to_all.py | 24 ++++++++++++++---------- 1 file changed, 14 insertions(+), 10 deletions(-) diff --git a/tests/utils_all_to_all.py b/tests/utils_all_to_all.py index 75a57ee2f..12c168815 100644 --- a/tests/utils_all_to_all.py +++ b/tests/utils_all_to_all.py @@ -175,14 +175,16 @@ async def _create_all_endpoints(self): if remote_port == self.listener.port: continue - ep = await self._create_endpoint(self.listener_address, remote_port) - self.bytes_bandwidth[remote_port] = [] - self.conns[(remote_port, i)] = ep + 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_port, ep, cache_only=True + self.listener.port, remote_address, ep, cache_only=True ) ) else: @@ -271,10 +273,12 @@ async def run_monitor(self): # Wait for all workers to connect while len(self.get_connections()) != self.num_workers: - print( - f"Waiting for all workers to connect, {len(self.get_connections())} " - f"of {self.num_workers} workers connected." - ) + 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 @@ -303,7 +307,7 @@ async def run_monitor(self): await self._close_connections_and_listener() def get_results(self): - for remote_port, bb in self.bytes_bandwidth.items(): + 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)) @@ -313,7 +317,7 @@ def get_results(self): "median bandwidth: %s/s" % ( ":".join([str(i) for i in local_address]), - ":".join([str(i) for i in remote_port]), + ":".join([str(i) for i in remote_address]), format_bytes(total_bytes), format_bytes(avg_bandwidth), format_bytes(median_bandwidth), From b1b0b1b948d0b51a5637916d0b69d24ac3ed84b3 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Mon, 12 Jul 2021 14:04:05 -0700 Subject: [PATCH 36/42] Remove usage of partial This prevets asyncio.iscoroutinefunction from returning False in Python < 3.8 --- tests/utils_comm_libs.py | 19 ++++++++----------- 1 file changed, 8 insertions(+), 11 deletions(-) diff --git a/tests/utils_comm_libs.py b/tests/utils_comm_libs.py index c7bdafa9f..b7dbdd0a8 100644 --- a/tests/utils_comm_libs.py +++ b/tests/utils_comm_libs.py @@ -1,5 +1,4 @@ import asyncio -import functools import struct import numpy as np @@ -217,13 +216,13 @@ def __init__(self, server, connections): @classmethod async def start_server(cls, host, port): - def _server_callback(connections, reader, writer): - connections.append(AsyncioCommConnection(reader, writer)) - connections = [] + def _server_callback(reader, writer): + connections.append(AsyncioCommConnection(reader, writer)) + server = await asyncio.start_server( - functools.partial(_server_callback, connections), host, port, limit=2 ** 30, + _server_callback, host, port, limit=2 ** 30, ) return cls(server, connections) @@ -299,16 +298,14 @@ def __init__(self, server, connections): @classmethod async def start_server(cls, host, port, listener_func): - async def _server_callback(connections, ep): + connections = [] + + async def _server_callback(ep): conn = UCXConnection(ep) connections.append(conn) await listener_func(conn) - connections = [] - - server = ucp.create_listener( - functools.partial(_server_callback, connections), port, - ) + server = ucp.create_listener(_server_callback, port) return cls(server, connections) def get_connections(self): From d3e835d4733b187e39d9c87df6eadc1aac59a035 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Tue, 13 Jul 2021 04:23:09 -0700 Subject: [PATCH 37/42] Add all-to-all benchmark with uvloop --- tests/test_multiple_processes_all_to_all.py | 9 +++++++-- tests/utils_all_to_all.py | 14 +++++++++++++- 2 files changed, 20 insertions(+), 3 deletions(-) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index dcfdec88e..ed2a9e447 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -1,7 +1,12 @@ import multiprocessing import pytest -from utils_all_to_all import asyncio_process, tornado_process, ucx_process +from utils_all_to_all import ( + asyncio_process, + tornado_process, + ucx_process, + uvloop_process, +) import ucp @@ -81,7 +86,7 @@ def _test_send_recv_cu( @pytest.mark.parametrize("size", [2 ** 20]) @pytest.mark.parametrize("iterations", [5]) @pytest.mark.parametrize( - "communication", [ucx_process, asyncio_process, tornado_process] + "communication", [ucx_process, asyncio_process, uvloop_process, tornado_process] ) def 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 index 12c168815..547627229 100644 --- a/tests/utils_all_to_all.py +++ b/tests/utils_all_to_all.py @@ -457,6 +457,7 @@ def asyncio_process( signal=None, ports=None, lock=None, + loop=None, ): w = AsyncioProcess( listener_address, @@ -472,10 +473,21 @@ def asyncio_process( lock=lock, ) run_func = w.run_monitor if is_monitor else w.run_worker - asyncio.get_event_loop().run_until_complete(run_func()) + 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, From bd86b1198be0116fcf4ab3ccbf5482e6bd5a63ff Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Tue, 13 Jul 2021 04:28:15 -0700 Subject: [PATCH 38/42] Add uvloop support for multi-node benchmark --- tests/benchmark_multiple_processes_all_to_all.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/tests/benchmark_multiple_processes_all_to_all.py b/tests/benchmark_multiple_processes_all_to_all.py index 77c3ec4c8..32e466ba9 100644 --- a/tests/benchmark_multiple_processes_all_to_all.py +++ b/tests/benchmark_multiple_processes_all_to_all.py @@ -1,7 +1,12 @@ import argparse import multiprocessing -from utils_all_to_all import asyncio_process, tornado_process, ucx_process +from utils_all_to_all import ( + asyncio_process, + tornado_process, + ucx_process, + uvloop_process, +) import ucp @@ -77,7 +82,7 @@ def parse_args(): default="ucx", type=str, help="Communication library to benchmark. Options are " - "'ucx', 'asyncio', 'tornado'. Default: 'ucx'.", + "'ucx', 'asyncio', 'tornado', 'uvloop'. Default: 'ucx'.", ) parser.add_argument( "--size", @@ -123,6 +128,8 @@ def main(): 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: From a96358ce6e668dd0074f71e16d93e31b805093d0 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Tue, 13 Jul 2021 04:57:00 -0700 Subject: [PATCH 39/42] Fix bandwidth calculation --- tests/utils_all_to_all.py | 14 ++++++++------ 1 file changed, 8 insertions(+), 6 deletions(-) diff --git a/tests/utils_all_to_all.py b/tests/utils_all_to_all.py index 547627229..063b5514b 100644 --- a/tests/utils_all_to_all.py +++ b/tests/utils_all_to_all.py @@ -72,7 +72,9 @@ 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) @@ -84,6 +86,8 @@ async def _transfer(self, ep, msg2send, send_first=True): 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) @@ -92,15 +96,13 @@ async def _listener(self, ep): 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) * self.iterations + send_recv_bytes = nbytes(message) * 2 - t = monotonic() - await self._transfer(ep, message, send_first=False) - total_time = monotonic() - t + time_per_iteration = await self._transfer(ep, message, send_first=False) if cache_only is False: - self.bytes_bandwidth[worker_address].append( - (send_recv_bytes, send_recv_bytes / total_time) + self.bytes_bandwidth[worker_address] += list( + (send_recv_bytes, send_recv_bytes / t) for t in time_per_iteration ) def _init(self): From b3bdfb0d3f00fc79ec7cbdbe30aac71848e514ab Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Tue, 13 Jul 2021 06:21:52 -0700 Subject: [PATCH 40/42] Parse bytes with --size --- tests/benchmark_multiple_processes_all_to_all.py | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/tests/benchmark_multiple_processes_all_to_all.py b/tests/benchmark_multiple_processes_all_to_all.py index 32e466ba9..aee6404f9 100644 --- a/tests/benchmark_multiple_processes_all_to_all.py +++ b/tests/benchmark_multiple_processes_all_to_all.py @@ -8,6 +8,8 @@ uvloop_process, ) +from dask.utils import parse_bytes + import ucp @@ -46,6 +48,7 @@ def parse_args(): ) parser.add_argument( "--monitor-address", + metavar="IP:PORT", default=None, help="Address where --monitor process is listening to in the HOST:PORT " "format.", @@ -86,9 +89,10 @@ def parse_args(): ) parser.add_argument( "--size", - default=2 ** 20, - type=int, - help="Size to be passed for data generation function. Default: " "1048576.", + metavar="BYTES", + default="1 MiB", + type=parse_bytes, + help="Size to be passed for data generation function. Default: '1 Mb'.", ) parser.add_argument( "--iterations", From 29468b49dabdc6a5f3a28a1fb5f905af5d35b8e5 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Tue, 13 Jul 2021 06:26:04 -0700 Subject: [PATCH 41/42] Skip uvloop test when not installed --- tests/test_multiple_processes_all_to_all.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/test_multiple_processes_all_to_all.py b/tests/test_multiple_processes_all_to_all.py index ed2a9e447..6f139a7fc 100644 --- a/tests/test_multiple_processes_all_to_all.py +++ b/tests/test_multiple_processes_all_to_all.py @@ -91,6 +91,9 @@ def _test_send_recv_cu( 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, From 14010ac47b0961c0798a64c0696ee0a9acaca599 Mon Sep 17 00:00:00 2001 From: Peter Andreas Entschev Date: Fri, 6 Aug 2021 07:56:09 -0700 Subject: [PATCH 42/42] Raise when number of workers don't match --- tests/utils_all_to_all.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/tests/utils_all_to_all.py b/tests/utils_all_to_all.py index 063b5514b..892c70704 100644 --- a/tests/utils_all_to_all.py +++ b/tests/utils_all_to_all.py @@ -152,7 +152,13 @@ async def _wait_for_workers(self): _, worker_addresses = await self.monitor_ep.recv() assert worker_addresses["data"][0] == OP_CLUSTER_READY self.worker_addresses = worker_addresses["data"][1] - assert len(self.worker_addresses) == self.num_workers + 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: