Numba并行JIT函数访问全局变量时出现计算异常的问题咨询
Numba并行JIT函数访问全局变量时出现计算异常的问题咨询
大家好,我最近在使用Numba时遇到了一个棘手的问题,想请教一下社区的朋友们。
首先我知道,常规的Numba JIT函数没办法动态访问全局变量——变量值在编译时就被固定了,但可以通过objmode强制实现Python风格的动态访问,当然这样会带来一定的性能损耗,示例代码如下:
from numba import njit, objmode @njit def access_global_variable_1(): return global_variable @njit def access_global_variable_2(): with objmode(retval='int64'): retval = globals()['global_variable'] return retval global_variable = 0 access_global_variable_1() # returns 0 access_global_variable_2() # returns 0 global_variable = 1 access_global_variable_1() # returns 0 access_global_variable_2() # returns 1
这个示例里,access_global_variable_2借助objmode确实能动态获取全局变量的最新值,一切正常。
但当我给函数加上并行化逻辑后,问题就出现了。我写了几个测试函数来验证:
from numba import njit, objmode, prange import numpy as np @njit def _add(a): with objmode(): globals()['acc'] += a @njit(parallel=True) def parallel_sum_global(arr): for i in prange(len(arr)): _add(arr[i]) @njit(parallel=False) def sum_global(arr): for i in prange(len(arr)): _add(arr[i]) @njit(parallel=True) def parallel_sum_local(arr): acc = 0 for i in prange(len(arr)): acc += arr[i] return acc n = 100 print('True answer:', np.arange(n).sum()) # True answer: 4950 acc = 0 parallel_sum_global(np.arange(n)) print('Numba parallel global answer:', acc) # Numba parallel answer: 78 acc = 0 sum_global(np.arange(n)) print('Numba global answer:', acc) # Numba global answer: 4950 acc = parallel_sum_local(np.arange(n)) print('Numba parallel local answer:', acc) # Numba parallel local answer: 4950
运行结果非常反常:
- 非并行的
sum_global能正确得到结果4950 - 使用局部变量的并行函数
parallel_sum_local也能输出正确结果 - 但同时使用全局变量和并行化的
parallel_sum_global,结果却只有78,完全偏离预期
我还注意到一个细节:parallel_sum_global似乎只正确处理了数组的前12个元素,后续的元素都没有被正确累加。我当前使用的是8核的开发电脑。
这个问题的触发条件很明确:只有同时结合全局变量访问和并行化时才会出现,单独使用其中一种逻辑都不会有问题。有没有朋友遇到过类似的情况?或者有没有可行的解决办法呢?
(注:这个是我简化出来的最小复现示例,实际是在更复杂的业务代码中遇到的这个问题)
备注:内容来源于stack exchange,提问作者Hmwat
相关产品推荐
相关产品推荐

