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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
104 changes: 99 additions & 5 deletions src/executorlib/standalone/interactive/communication.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@

import cloudpickle
import zmq
import zmq.asyncio


class ExecutorlibSocketError(RuntimeError):
Expand Down Expand Up @@ -32,7 +33,7 @@ def __init__(
log_obj_size (boolean): Enable debug mode which reports the size of the communicated objects.
time_out_ms (int): Time out for waiting for a message on socket in milliseconds.
"""
self._context = zmq.Context()
self._context = self._new_context()
self._socket = self._context.socket(zmq.PAIR)
self._poller = zmq.Poller()
self._poller.register(self._socket, zmq.POLLIN)
Expand All @@ -46,6 +47,9 @@ def __init__(
self._booted_sucessfully: bool = False
self._stop_function: Optional[Callable] = None

def _new_context(self) -> zmq.Context:
return zmq.Context()

@property
def status(self) -> bool:
return self._booted_sucessfully
Expand Down Expand Up @@ -179,13 +183,95 @@ def __del__(self):
self.shutdown(wait=True)


class AsyncSocketInterface(SocketInterface):
"""
Variant of the SocketInterface based on zmq.asyncio. The communication methods are coroutines which have to be
awaited on the event loop the interface is used from. All interfaces of one executor share a single ZMQ context,
which is owned and terminated by the caller. The communication protocol is identical to the SocketInterface.

Args:
context (zmq.asyncio.Context): shared asyncio ZMQ context
spawner (executorlib.shared.spawner.BaseSpawner): Interface for starting the parallel process
log_obj_size (boolean): Enable debug mode which reports the size of the communicated objects.
time_out_ms (int): Time out for waiting for a message on socket in milliseconds.
"""

def __init__(
self,
context: zmq.asyncio.Context,
spawner=None,
log_obj_size: bool = False,
time_out_ms: int = 1000,
):
self._shared_context = context
super().__init__(
spawner=spawner, log_obj_size=log_obj_size, time_out_ms=time_out_ms
)

def _new_context(self) -> zmq.Context:
return self._shared_context

async def send_dict_async(self, input_dict: dict):
data = cloudpickle.dumps(input_dict)
if self._logger is not None:
self._logger.warning("Send dictionary of size: " + str(sys.getsizeof(data)))
await self._socket.send(data)

async def receive_dict_async(self) -> dict:
response = 0
while response == 0:
response = await self._socket.poll(self._time_out_ms)
if not self._spawner.poll():
return {
"error": ExecutorlibSocketError(
"SocketInterface crashed during execution."
),
}
data = await self._socket.recv()
if self._logger is not None:
self._logger.warning(
"Received dictionary of size: " + str(sys.getsizeof(data))
)
return cloudpickle.loads(data)

async def send_and_receive_dict_async(self, input_dict: dict) -> dict:
await self.send_dict_async(input_dict=input_dict)
return await self.receive_dict_async()

async def shutdown_async(self, wait: bool = True):
result = None
if self._spawner.poll():
output = await self.send_and_receive_dict_async(
input_dict={"shutdown": True, "wait": wait}
)
if "result" in output:
result = output["result"]
self._spawner.shutdown(wait=wait)
self._reset_socket()
return result

def _reset_socket(self):
# The shared context is terminated by its owner, only the socket is closed here.
if self._socket is not None:
self._socket.close(linger=0)
self._process = None
self._socket = None
self._context = None

def __del__(self):
# The socket belongs to the event loop thread, so only the worker process is stopped here.
if self._spawner is not None and self._spawner.poll():
self._spawner.shutdown(wait=False)


def interface_bootup(
command_lst: list[str],
connections,
hostname_localhost: Optional[bool] = None,
log_obj_size: bool = False,
worker_id: Optional[int] = None,
stop_function: Optional[Callable] = None,
context: Optional[zmq.asyncio.Context] = None,
) -> SocketInterface:
"""
Start interface for ZMQ communication
Expand All @@ -205,6 +291,7 @@ def interface_bootup(
worker_id (int): Communicate the worker which ID was assigned to it for future reference and resource
distribution.
stop_function (Callable): Function to stop the interface.
context (zmq.asyncio.Context): if provided an AsyncSocketInterface using this shared context is returned.

Returns:
executorlib.shared.communication.SocketInterface: socket interface for zmq communication
Expand All @@ -218,10 +305,17 @@ def interface_bootup(
]
if worker_id is not None:
command_lst += ["--worker-id", str(worker_id)]
interface = SocketInterface(
spawner=connections,
log_obj_size=log_obj_size,
)
if context is not None:
interface: SocketInterface = AsyncSocketInterface(
context=context,
spawner=connections,
log_obj_size=log_obj_size,
)
else:
interface = SocketInterface(
spawner=connections,
log_obj_size=log_obj_size,
)
command_lst += [
"--zmqport",
str(interface.bind_to_random_port()),
Expand Down
38 changes: 26 additions & 12 deletions src/executorlib/task_scheduler/interactive/blockallocation.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import os
import queue
import random
from concurrent.futures import Future
Expand All @@ -17,6 +18,10 @@
from executorlib.standalone.interactive.spawner import BaseSpawner, MpiExecSpawner
from executorlib.standalone.queue import cancel_items_in_queue, put_front
from executorlib.task_scheduler.base import TaskSchedulerBase, validate_resource_dict
from executorlib.task_scheduler.interactive.blockallocation_async import (
AsyncWorker,
AsyncWorkerPool,
)
from executorlib.task_scheduler.interactive.shared import (
execute_task_dict,
reset_task_dict,
Expand Down Expand Up @@ -61,6 +66,8 @@ class BlockAllocationTaskScheduler(TaskSchedulerBase):

"""

_async_pool: Optional[AsyncWorkerPool] = None

def __init__(
self,
max_workers: int = 1,
Expand All @@ -86,14 +93,22 @@ def __init__(
self._alive_workers_lock = Lock()
self._bootup_events = [Event() for _ in range(self._max_workers)]
self._bootup_events[0].set()
# Experimental: drive all workers from one private asyncio event loop instead of one thread per worker.
if os.environ.get("EXECUTORLIB_ASYNCIO", "0").lower() in ("1", "true"):
assert self._future_queue is not None
self._async_pool = AsyncWorkerPool(future_queue=self._future_queue)
self._set_process(
process=[
Thread(
target=_execute_multiple_tasks,
kwargs=self._worker_kwargs(worker_id),
)
for worker_id in range(self._max_workers)
],
process=[self._new_worker(worker_id) for worker_id in range(max_workers)],
)

def _new_worker(self, worker_id: int):
if self._async_pool is not None:
return AsyncWorker(
pool=self._async_pool, kwargs=self._worker_kwargs(worker_id)
)
return Thread(
target=_execute_multiple_tasks,
kwargs=self._worker_kwargs(worker_id),
)

def _worker_kwargs(self, worker_id: int) -> dict:
Expand Down Expand Up @@ -137,10 +152,7 @@ def max_workers(self, max_workers: int):
with self._alive_workers_lock:
self._alive_workers[0] += max_workers - old_max_workers
new_process_lst = [
Thread(
target=_execute_multiple_tasks,
kwargs=self._worker_kwargs(worker_id),
)
self._new_worker(worker_id)
for worker_id in range(old_max_workers, max_workers)
]
for process_instance in new_process_lst:
Expand Down Expand Up @@ -210,10 +222,12 @@ def shutdown(self, wait: bool = True, *, cancel_futures: bool = False):
for process in self._process:
process.join()
self._future_queue.join()
if self._async_pool is not None:
self._async_pool.close(wait=wait)
self._process = None
self._future_queue = None

def _set_process(self, process: list[Thread]): # type: ignore
def _set_process(self, process: list): # type: ignore
"""
Set the process for the executor.

Expand Down
Loading
Loading