How to send jobs from multiple client threads to a pool of multiprocessor workers and get results returned (python)?

Viewed 205

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)
0 Answers
Related