使用@njit合并NumPy数组报错问题及修复咨询
问题场景
下面是一个合并NumPy数组的示例:未添加@njit装饰器时代码正常运行,但加上@njit后就会抛出TypingError错误。
初始化数组代码
import numpy as np from numba import njit arr1 = np.array([1, 2, 3]) arr2 = np.array([4, 5, 6]) arr3 = np.array([7, 8, 9]) print('arr1 :', arr1) print() print('arr2 :', arr2) print() print('arr3 :', arr3)
未使用@njit的正常运行代码
def union_of_arrays(): arr = np.array([arr1, arr2, arr3]).T return arr arr = union_of_arrays() print(arr)
添加@njit后的报错代码
@njit def union_of_arrays(): arr = np.array([arr1, arr2, arr3]).T return arr arr = union_of_arrays() print(arr)
报错信息
TypingError Traceback (most recent call last)
Cell In[4], line 8
4 arr = np.array([arr1, arr2, arr3]).T
6 return arr
----> 8 arr = union_of_arrays()
9 print(arr)File c:\Users\user\AppData\Local\Programs\Python\Python310\lib\site-packages\numba\core\dispatcher.py:468, in _DispatcherBase._compile_for_args(self, *args, **kws)
464 msg = (f"{str(e).rstrip()}This error may have been caused "
465 f"by the following argument(s):
{args_str}
")
466 e.patch_message(msg)
--> 468 error_rewrite(e, 'typing')
469 except errors.UnsupportedError as e:
470 # Something unsupported is present in the user code, add help info
471 error_rewrite(e, 'unsupported_error')File c:\Users\user\AppData\Local\Programs\Python\Python310\lib\site-packages\numba\core\dispatcher.py:409, in _DispatcherBase._compile_for_args..error_rewrite(e, issue_type)
407 raise e
408 else:
--> 409 raise e.with_traceback(None)TypingError: Failed in nopython mode pipeline (step: nopython frontend)
No implementation of function Function() found for signature:...
File "....\AppData\Local\Temp\ipykernel_8316\2415862080.py", line 4:
修复方案
问题根源
Numba的@njit(nopython模式)不支持通过np.array([arr1, arr2, arr3])这种列表嵌套数组的方式构造新数组,这种写法无法被Numba的类型推导系统正确识别。
可行修复代码
方案1:用np.column_stack直接生成目标数组
np.column_stack会把输入的一维数组按列拼接,正好等价于原代码中先堆叠再转置的效果,而且是Numba支持的操作:
@njit def union_of_arrays(): arr = np.column_stack((arr1, arr2, arr3)) return arr arr = union_of_arrays() print(arr)
方案2:用np.vstack堆叠后转置
如果想保留原逻辑的堆叠+转置结构,可以用np.vstack替代np.array([...]),np.vstack同样被Numba支持:
@njit def union_of_arrays(): arr = np.vstack((arr1, arr2, arr3)).T return arr arr = union_of_arrays() print(arr)
注意点
传递给np.vstack或np.column_stack的参数要写成元组(用圆括号()包裹),而不是列表(方括号[]),这能让Numba更准确地推导类型。
内容的提问来源于stack exchange,提问作者S.G.Studios

