I need to compute an expensive to compute function many times and want to use all processor cores for it. That would be relatively simple if I had all function argument sets at once: I could use multiprocessing pool.map . However I don't have them all at once, and also would like to avoid starting a separate process for each function computation. Therefore I want to start a pool of workers (no problems), send them the jobs from the clients (use queue, no problems), and then get the results back to the client (but how?).
More specifically, I need to compute my function on a multi-dimensional mesh array of arguments. In my case it is possible to avoid the computation for every array point by using two specific propeties of my function: it is monotonously nondecreasing for all arguments, and it can take only a few discrete values. Thus, if the function value is equal for two separate points on the mesh, it will be the same for all points in-between. Now we divide our mesh into parts and use recursion. Using recursion however means I cannot easily use pool.map
Actually I found a working solution, but guess there could be a more straightforward and reliable way.
I set up a separate dispatcher thread. It gets ALL results from the workers through a results queue, and then lets clients to pick up their results. Below the code for a simplified (2d) case.
EDIT: to the comment by Shihab Shahriar. Large part of the code needed to make the whole work, but do not directly relate to the question. Specifically:
contpl(), sample_f() are auxillary functions. In the real problem, instead of sample_f() there would be a complex and expensive simulation.
recfill() is the recursive function. It needs the resuls from the expensive calculation, and receives them via calling getz() . It would recursively start instances of itself in separate threads. Other details here unimportant.
getz() is an important part of my solution: It is a proxy between the above recursive function, and the workers pool. It sends the job parameters (tagged by id) via the queue task_q to the workers pool, then waits for events from the dispatcher() and checks if the calculated results already arrived, then returns them to it's caller. As the parent recfill() instances run in multiple threads, getz() too.
Worker() - instances of the class run in separate processes, wait for jobs arriving from the queue task_q, call the "expensive function", and put the results together with the id tag into the result_queue / results_q
dispatcher() runs in a separate thread, receives the results from results_q and puts them into a shared dict with the id as index. Then sends an event to getz() instances to check, whose result arrived.
main - starts workers, starts dispatcher, calls the recfill() , and cleans up.
from concurrent.futures import ThreadPoolExecutor as Pool
import multiprocessing
import threading
import numpy as np
import matplotlib.pyplot as plt
def contpl(p1, p2, Z):
"""
plots color density plot of XYZf
"""
plt.figure(figsize=(5,5))
plt.contourf(p1, p2, Z, 3, cmap='RdYlGn')
plt.show()
def sample_f(x, y):
"""
simple sample monotonous function to plot quater-circles
returns integer values from 0 to 2
"""
return np.round(1.4 * np.sqrt(x**2 + y**2))
def getz(ix, iy, mp_params):
"""
gets the values of the function sample_f for (x, y) values
from the parameter arrays p1, p2 for the indices ix, iy
if the value has not been already computed,
send the job to a worker, then wait until the result is ready
"""
task_q = mp_params["task_q"]
result_event = mp_params["result_event"]
result_dict = mp_params["result_dict"]
num_workers = mp_params["num_workers"]
results_q = mp_params["results_q"]
z = Zarr[ix, iy]
if z >= 0:
return z # nice, the point has already been calculated
# otherwise z is -1 from the array initialisazion
else:
# compute "flattened index" of a point as id
dims = Zarr.shape
id = np.ravel_multi_index((ix, iy), dims)
task_q.put((id, p1[ix, iy], p2[ix, iy])) # send the job, targeted by the id, to workers
# not wait until dispatcher calls
while True:
result_event.wait()
try:
# anything for me?
z = result_dict.pop(id)
result_event.clear()
break
except KeyError:
pass
# now the point is computed, write the value into the array
Zarr[ix, iy] = z
return z
def recfill(ix, iy, mp_params):
"""
recursive function to compute values of a monotonous function
on a 2D square (sub-)array of parameters
"""
(ix0, ix1) = ix # x indices
(iy0, iy1) = iy # y indices
z0 = getz(ix0, iy0, mp_params) # get the bottom left point
# if the array size is one in all dimensions, we reached the recursion limit
if (ix0 == ix1) and (iy0 == iy1):
return
else:
# get the top right point
z1 = getz(ix1, iy1, mp_params)
# if the values for bottom left and top right are equal, they are the same for all
# elements in between
if z0 == z1:
Zarr[ix0:ix1+1, iy0:iy1+1] = z0 # fill in the subarray
return # and we are done for this recursion branch
else:
# divide the sub-array by half in each dimension
xhalf = (ix1 - ix0 + 1) // 2
yhalf = (iy1 - iy0 + 1) // 2
ixlo = (ix0, ix0+xhalf-1)
iylo = (iy0, iy0+yhalf-1)
ixhi = (ix0+xhalf, ix1)
iyhi = (iy0+yhalf, iy1)
# prepare arguments for the map function
l1 = [(ixlo, iylo), (ixlo, iyhi), (ixhi, iylo), (ixhi, iyhi)]
(ixs, iys) = zip(*l1)
mpps = [mp_params]*4
# and now multithreaded recursive call for each quater of the initial sub-array
with Pool() as p:
p.map(recfill, ixs, iys, mpps)
return
class Worker(multiprocessing.Process):
"""
adapted from
https://pymotw.com/3/multiprocessing/communication.html
"""
def __init__(self, mp_params):
multiprocessing.Process.__init__(self)
self.task_queue = mp_params["task_q"]
self.result_queue = mp_params["results_q"]
def run(self):
proc_name = self.name
while True:
job = self.task_queue.get()
if job is None:
print('{}: Exiting'.format(proc_name))
break
(id, x, y) = job
result = sample_f(x, y)
answer = (id, result)
self.result_queue.put(answer)
def dispatcher(mp_params):
"""
receives the computation results from the results queue,
puts them into a shared dictionary,
and notifies all clients per event,
that they should check the dictionary,
if there is anything for them
"""
result_event = mp_params["result_event"]
result_dict = mp_params["result_dict"]
results_q = mp_params["results_q"]
while True:
qitem = results_q.get()
if qitem is not None:
(id, result) = qitem
result_dict[id] = result
result_event.set()
else:
break
if __name__ == '__main__':
result_event = threading.Event()
num_workers = multiprocessing.cpu_count()
task_q = multiprocessing.SimpleQueue()
results_q = multiprocessing.Queue() # why using SimpleQueue here would hang the program?
result_dict = {}
mp_params = {}
mp_params["task_q"] = task_q
mp_params["results_q"] = results_q
mp_params["result_dict"] = result_dict
mp_params["result_event"] = result_event
mp_params["num_workers"] = num_workers
print('Creating {} workers'.format(num_workers))
workers = [Worker(mp_params) for i in range(num_workers)]
for w in workers:
w.start()
# creating dispatcher thread
t = threading.Thread(target=dispatcher, args=(mp_params, ))
t.start()
# creating parameter arrays
arrsize = 128
xvec = np.linspace(0, 1, arrsize)
yvec = np.linspace(0, 1, arrsize)
(p1, p2) = np.meshgrid(xvec, yvec)
# initialize the results array
# our sample_f returns only non-negative values
# therefore fill in with -1 to indicate the values
# which have not been computed yet
Zarr = np.full_like(p1, -1, dtype=np.int8)
# now call our recursive function
# to compute all array values
recfill((0,arrsize-1), (0,arrsize-1), mp_params)
# clean up
for i in range(num_workers):
task_q.put(None) # stop all workers
results_q.put(None) # stop dispatcher
t.join()
# plot the results
contpl(p1, p2, Zarr)
# and check the results by comparing with directly
# calculated values
Z = sample_f(p1, p2)
assert np.all(Z == Zarr)