Numba jit装饰器num_threads参数报错及替代方案咨询
环境与需求
import numpy as np import numba as nb import multiprocessing import sklearn print(np.__version__,nb.__version__)# numpy: 1.23.5 numba: 0.56.4 np.random.seed(0) a=np.random.rand(int(6e+8)) # 超大一维numpy数组
需求:计算上述超大一维numpy数组的L2范数。使用sklearn.preprocessing.normalize(a.reshape(1,-1), norm="l2", axis=1)耗时14.5秒,尝试用Numba优化并指定线程数加速。
尝试的代码与报错
尝试使用numba.jit()装饰器并添加num_threads参数:
@nb.jit(nopython=True,nogil=True,parallel=True,fastmath=True,num_threads=multiprocessing.cpu_count()) def normalize_numba_optimized(array): sum_of_squares = 0.0 # 计算数组元素的平方和 for elem in array: sum_of_squares += elem * elem # 计算数组的范数 norm = np.sqrt(sum_of_squares) # 将数组除以范数 for i in nb.prange(array.size): array[i] /= norm return array
执行时出现报错:
KeyError Traceback (most recent call last) KeyError: "Unrecognized options: {'num_threads'}. Known options are dict_keys(['_dbg_extend_lifetimes', '_dbg_optnone', '_nrt', 'boundscheck', 'debug', 'error_model', 'fastmath', 'forceinline', 'forceobj', 'inline', 'looplift', 'no_cfunc_wrapper', 'no_cpython_wrapper', 'no_rewrites', 'nogil', 'nopython', 'parallel', 'target_backend'])"
移除num_threads参数后代码可正常运行,耗时3.97秒。
问题
num_threads是否已从numba.jit()的参数中废弃?或有无替代方案实现指定线程数的需求?
更新提示
经提示,可通过print(nb.get_num_threads())查看可用CPU核心数,再用nb.set_num_threads(multiprocessing.cpu_count())设置线程数。
解答
关于num_threads参数的说明
num_threads从来不是nb.jit()装饰器的合法参数,报错信息中已明确列出了jit支持的所有参数,其中并不包含该选项,因此不存在“废弃”一说,从始至终它就不属于jit的参数列表。
指定线程数的替代方案
有两种常用方式可以实现指定线程数的需求:
全局线程数设置
使用nb.set_num_threads(n)函数(n为目标线程数,例如multiprocessing.cpu_count())来全局设置Numba并行代码的线程数,该设置会影响后续所有的Numba并行执行逻辑。
可以用nb.get_num_threads()查看当前的线程数配置。针对单个并行循环设置
如果不需要全局修改线程数,可以在使用nb.prange()时,直接给它传递num_threads参数,指定当前循环的线程数:for i in nb.prange(array.size, num_threads=multiprocessing.cpu_count()): array[i] /= norm
额外优化建议
你当前代码中计算平方和的循环是串行执行的,可以改成并行计算进一步提升效率:
@nb.jit(nopython=True,nogil=True,parallel=True,fastmath=True) def normalize_numba_optimized(array): sum_of_squares = 0.0 # 并行计算平方和 for i in nb.prange(array.size): sum_of_squares += array[i] * array[i] norm = np.sqrt(sum_of_squares) for i in nb.prange(array.size): array[i] /= norm return array
Numba会自动处理并行累加的线程安全问题,拆分计算任务并合并结果,进一步缩短执行时间。
内容的提问来源于stack exchange,提问作者farid

