How to avoid recompilation of numba code when using multiprocessing

Viewed 83

I am new to numba, and I found that numba function is very fast when using signle processing (numpy 0.20s and numba 0.04s for running 1000 times), but it is very slow when using multi processing under jit, and it is not working under aot. I have tried nogil, cache, parallel and etc. My numpy or numba function is cpu intensive, so I cannot use thread or prange.

base.py

import numpy as np
from numba import jit
from numba.pycc import CC


cc = CC("my_numba")


class MathHelper:
    @staticmethod
    def np_ffill(arr: np.array):
        mask = np.isnan(arr)
        if arr.ndim == 1:
            idx = np.where(~mask, np.arange(mask.size), 0)
            np.maximum.accumulate(idx, out=idx)
            return arr[idx]
    @staticmethod
    @jit(nopython=True)
    def nb_ffill_jit_v1(arr: np.array):
        values = arr.copy()
        if arr.ndim == 1:
            last_value = np.nan
            for n in range(arr.size):
                if np.isnan(arr[n]):
                    values[n] = last_value
                else:
                    last_value = arr[n]
        return values
    @staticmethod
    @jit(nopython=True, nogil=True)
    def nb_ffill_jit_v2(arr: np.array):
        ...  # (copy nb_ffill_jit_v1 code to here)
    @staticmethod
    @jit(nopython=True, nogil=True, cache=True)
    def nb_ffill_jit_v3(arr: np.array):
        ...
    @staticmethod
    @cc.export("nb_ffill_aot_v1", "f8[:](f8[:])")
    def nb_ffill_aot_v1(arr: np.array):
        ...
    @staticmethod
    @jit(nopython=True, parallel=True)
    @cc.export("nb_ffill_aot_v2", "f8[:](f8[:])")
    def nb_ffill_aot_v2(arr: np.array):
        ...
    @staticmethod
    @jit(nopython=True, parallel=True, nogil=True)
    @cc.export("nb_ffill_aot_v3", "f8[:](f8[:])")
    def nb_ffill_aot_v3(arr: np.array):
        ...
    @staticmethod
    @jit(nopython=True, parallel=True, nogil=True, cache=True)
    @cc.export("nb_ffill_aot_v4", "f8[:](f8[:])")
    def nb_ffill_aot_v4(arr: np.array):
        ...


if __name__ == "__main__":
    cc.compile()

numba_test.py

import concurrent.futures
import time
from typing import Union

import numpy as np

import my_numba
from base import MathHelper


def np_arr(n: int, low: Union[int, float] = 1000, high: Union[int, float] = 2000, digit: int = 0):
    return np.around(np.random.rand(n) * (high - low) + low, digit)

def np_arr_nan(n: int, n_nan: int, low: Union[int, float] = 1000, high: Union[int, float] = 2000, digit: int = 0):
    idx = np.random.choice(n, n_nan, replace=False)
    arr = np_arr(n, low, high, digit)
    arr[idx] = np.nan
    return arr

def multi_timeit(func, *args):
    time_start = time.time()
    with concurrent.futures.ProcessPoolExecutor(max_workers=25) as executor:
        for _ in range(25):
            executor.submit(func, args)
    return time.time() - time_start

def single_timeit(func, *args):
    time_start = time.time()
    for _ in range(1000):
        func(*args)
    return time.time() - time_start

if __name__ == "__main__":
    arr_nan_small = np_arr_nan(40000, 1000, digit=2)
    
    MathHelper.nb_ffill_jit_v1(arr_nan_small)
    print(single_timeit(MathHelper.np_ffill, arr_nan_small))
    print(single_timeit(MathHelper.nb_ffill_jit_v1, arr_nan_small))

    print(multi_timeit(MathHelper.np_ffill, arr_nan_small))
    print(multi_timeit(MathHelper.nb_ffill_jit_v1, arr_nan_small))
    print(multi_timeit(MathHelper.nb_ffill_jit_v2, arr_nan_small))
    print(multi_timeit(MathHelper.nb_ffill_jit_v3, arr_nan_small))
    print(multi_timeit(MathHelper.nb_ffill_aot_v1, arr_nan_small))
    print(multi_timeit(MathHelper.nb_ffill_aot_v2, arr_nan_small))
    print(multi_timeit(MathHelper.nb_ffill_aot_v3, arr_nan_small))
    print(multi_timeit(MathHelper.nb_ffill_aot_v4, arr_nan_small))
    print(multi_timeit(my_numba.nb_ffill_aot_v1, arr_nan_small))  # not working
    print(multi_timeit(my_numba.nb_ffill_aot_v2, arr_nan_small))  # not working
    print(multi_timeit(my_numba.nb_ffill_aot_v3, arr_nan_small))  # not working
    print(multi_timeit(my_numba.nb_ffill_aot_v4, arr_nan_small))  # not working
0 Answers
Related