(
func: callable, data, n_proc, target_data_type="ndarray", cpu_intensive=True, use_worker_id=False
)
| 106 | |
| 107 | |
| 108 | def parallel_data_prefetch( |
| 109 | func: callable, data, n_proc, target_data_type="ndarray", cpu_intensive=True, use_worker_id=False |
| 110 | ): |
| 111 | # if target_data_type not in ["ndarray", "list"]: |
| 112 | # raise ValueError( |
| 113 | # "Data, which is passed to parallel_data_prefetch has to be either of type list or ndarray." |
| 114 | # ) |
| 115 | if isinstance(data, np.ndarray) and target_data_type == "list": |
| 116 | raise ValueError("list expected but function got ndarray.") |
| 117 | elif isinstance(data, abc.Iterable): |
| 118 | if isinstance(data, dict): |
| 119 | print( |
| 120 | f'WARNING:"data" argument passed to parallel_data_prefetch is a dict: Using only its values and disregarding keys.' |
| 121 | ) |
| 122 | data = list(data.values()) |
| 123 | if target_data_type == "ndarray": |
| 124 | data = np.asarray(data) |
| 125 | else: |
| 126 | data = list(data) |
| 127 | else: |
| 128 | raise TypeError( |
| 129 | f"The data, that shall be processed parallel has to be either an np.ndarray or an Iterable, but is actually {type(data)}." |
| 130 | ) |
| 131 | |
| 132 | if cpu_intensive: |
| 133 | Q = mp.Queue(1000) |
| 134 | proc = mp.Process |
| 135 | else: |
| 136 | Q = Queue(1000) |
| 137 | proc = Thread |
| 138 | # spawn processes |
| 139 | if target_data_type == "ndarray": |
| 140 | arguments = [ |
| 141 | [func, Q, part, i, use_worker_id] |
| 142 | for i, part in enumerate(np.array_split(data, n_proc)) |
| 143 | ] |
| 144 | else: |
| 145 | step = ( |
| 146 | int(len(data) / n_proc + 1) |
| 147 | if len(data) % n_proc != 0 |
| 148 | else int(len(data) / n_proc) |
| 149 | ) |
| 150 | arguments = [ |
| 151 | [func, Q, part, i, use_worker_id] |
| 152 | for i, part in enumerate( |
| 153 | [data[i: i + step] for i in range(0, len(data), step)] |
| 154 | ) |
| 155 | ] |
| 156 | processes = [] |
| 157 | for i in range(n_proc): |
| 158 | p = proc(target=_do_parallel_data_prefetch, args=arguments[i]) |
| 159 | processes += [p] |
| 160 | |
| 161 | # start processes |
| 162 | print(f"Start prefetching...") |
| 163 | import time |
| 164 | |
| 165 | start = time.time() |
no outgoing calls
no test coverage detected