You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Numba jit装饰器num_threads参数报错及替代方案咨询

Numba指定线程数问题解答

环境与需求

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的参数列表。

指定线程数的替代方案

有两种常用方式可以实现指定线程数的需求:

  1. 全局线程数设置
    使用nb.set_num_threads(n)函数(n为目标线程数,例如multiprocessing.cpu_count())来全局设置Numba并行代码的线程数,该设置会影响后续所有的Numba并行执行逻辑。
    可以用nb.get_num_threads()查看当前的线程数配置。

  2. 针对单个并行循环设置
    如果不需要全局修改线程数,可以在使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.07 06:49:53