|
| 1 | +import multiprocessing as mp |
| 2 | +from enum import Enum |
| 3 | +from typing import Callable, List, Optional |
| 4 | + |
| 5 | +import tqdm |
| 6 | + |
| 7 | +try: |
| 8 | + from dask.distributed import Client |
| 9 | + |
| 10 | + dask_available = True |
| 11 | +except ImportError: |
| 12 | + dask_available = False |
| 13 | + |
| 14 | + |
| 15 | +def serial_execution(func: Callable, entries: List): |
| 16 | + return [func(args) for args in tqdm.tqdm(entries, total=len(entries))] |
| 17 | + |
| 18 | + |
| 19 | +def parallel_execution_multiprocessing(func: Callable, entries: List, process_count: int): |
| 20 | + pool = mp.Pool(process_count) |
| 21 | + results = [entry for entry in tqdm.tqdm(pool.imap(func, entries), total=len(entries))] |
| 22 | + pool.close() |
| 23 | + |
| 24 | + return results |
| 25 | + |
| 26 | + |
| 27 | +if dask_available: |
| 28 | + def parallel_execution_dask_local(func: Callable, entries: List, client: Client): |
| 29 | + return client.gather(client.map(func, entries)) |
| 30 | + |
| 31 | + def parallel_execution_dask_cluster(func: Callable, entries: List, client: Client): |
| 32 | + return client.gather(client.map(func, entries)) |
| 33 | + |
| 34 | + |
| 35 | +class RunnerMode(Enum): |
| 36 | + SERIAL = "serial" |
| 37 | + MULTIPROCESSING = "multiprocessing" |
| 38 | + DASK_LOCAL = "dask_local" |
| 39 | + DASK_JOB_QUEUE_CLUSTER = "dask_job_queue_cluster" |
| 40 | + |
| 41 | + |
| 42 | +class Runner: |
| 43 | + mode: RunnerMode |
| 44 | + thread_count: Optional[int] |
| 45 | + process_count: Optional[int] |
| 46 | + dask_cluster = None |
| 47 | + client = None |
| 48 | + |
| 49 | + def __init__(self, mode: RunnerMode, thread_count: Optional[int] = None, process_count: Optional[int] = None, dask_cluster = None): |
| 50 | + if mode in [RunnerMode.DASK_LOCAL, RunnerMode.DASK_JOB_QUEUE_CLUSTER]: |
| 51 | + assert dask_available is not None, "Execution using Dask requires installation of optional dependencies. The optional pip package group is called 'cluster'" |
| 52 | + |
| 53 | + if mode is RunnerMode.SERIAL: |
| 54 | + assert thread_count is None and process_count is None and dask_cluster is None, "Serial execution doesn't take any parameters." |
| 55 | + elif mode in [RunnerMode.MULTIPROCESSING]: |
| 56 | + assert thread_count is None and process_count is not None and dask_cluster is None, "Only process count is needed." |
| 57 | + elif mode in [RunnerMode.DASK_LOCAL]: |
| 58 | + assert (thread_count is not None or process_count is not None) and dask_cluster is None, "Only process count and/or thread count are needed." |
| 59 | + elif mode in [RunnerMode.DASK_JOB_QUEUE_CLUSTER]: |
| 60 | + assert thread_count is None and process_count is None and dask_cluster is not None, "Dask execution takes only a Dask cluster object." |
| 61 | + else: |
| 62 | + assert False |
| 63 | + |
| 64 | + self.mode = mode |
| 65 | + self.thread_count = thread_count |
| 66 | + self.process_count = process_count |
| 67 | + self.cluster = dask_cluster |
| 68 | + |
| 69 | + if mode in [RunnerMode.DASK_LOCAL]: |
| 70 | + self.client = Client(n_workers=process_count, threads_per_worker=thread_count) |
| 71 | + |
| 72 | + if mode in [RunnerMode.DASK_JOB_QUEUE_CLUSTER]: |
| 73 | + self.client = Client(dask_cluster) |
| 74 | + |
| 75 | + def run(self, func: Callable, entries: List): |
| 76 | + if self.mode == RunnerMode.SERIAL: |
| 77 | + return serial_execution(func, entries) |
| 78 | + elif self.mode == RunnerMode.MULTIPROCESSING: |
| 79 | + return parallel_execution_multiprocessing(func, entries, process_count=self.process_count) |
| 80 | + elif self.mode == RunnerMode.DASK_LOCAL: |
| 81 | + return parallel_execution_dask_local(func, entries, self.client) |
| 82 | + if self.mode == RunnerMode.DASK_JOB_QUEUE_CLUSTER: |
| 83 | + return parallel_execution_dask_cluster(func, entries, self.client) |
| 84 | + else: |
| 85 | + assert False |
0 commit comments