使用numba调用numpy.dot报错的问题求助
问题:Numba装饰器下np.dot报错的解决
问题背景
使用@numba.njit()装饰器后,调用np.dot()处理二维int64数组时出现报错,移除装饰器则运行正常。尝试添加out=None参数也无效,期望得到和无Numba时相同的运行结果。
报错代码
import numpy as np import numba @numba.njit() def tst_dot(): a = np.array([[1, 0], [0, 1]]) b = np.array([[4, 1], [2, 2]]) return np.dot(a, b) print(tst_dot())
错误信息
No implementation of function Function(<function dot at 0x00000280CC542EF0>) found for signature: >>> dot(array(int64, 2d, C), array(int64, 2d, C)) There are 4 candidate implementations: - Of which 2 did not match due to: Overload in function 'dot_2': File: numba\np\linalg.py: Line 525. With argument(s): '(array(int64, 2d, C), array(int64, 2d, C))': Rejected as the implementation raised a specific error: TypingError: Failed in nopython mode pipeline (step: native lowering) Failed in nopython mode pipeline (step: nopython frontend) No implementation of function Function(<function dot at 0x00000280CC542EF0>) found for signature: >>> dot(array(int64, 2d, C), array(int64, 2d, C), array(int64, 2d, C)) There are 4 candidate implementations: - Of which 2 did not match due to: Overload in function 'dot_2': File: numba\np\linalg.py: Line 525. With argument(s): '(array(int64, 2d, C), array(int64, 2d, C), array(int64, 2d, C))': Rejected as the implementation raised a specific error: TypingError: too many positional arguments raised from C:\Users\a_che\PycharmProjects\minCovTarget\venv\lib\site-packages\numba\core\typing\templates.py:784 - Of which 2 did not match due to: Overload in function 'dot_3': File: numba\np\linalg.py: Line 784. With argument(s): '(array(int64, 2d, C), array(int64, 2d, C), array(int64, 2d, C))': Rejected as the implementation raised a specific error: LoweringError: Failed in nopython mode pipeline (step: native lowering) unsupported dtype for <BLAS function>() File "venv\lib\site-packages\numba\np\linalg.py", line 817: def codegen(context, builder, sig, args): <source elided> return lambda left, right, out: _impl(left, right, out) ^ During: lowering "$10call_function.4 = call $2load_deref.0(left, right, out, func=$2load_deref.0, args=[Var(left, linalg.py:817), Var(right, linalg.py:817), Var(out, linalg.py:817)], kws=(), vararg=None, varkwarg=None, target=None)" at C:\Users\a_che\PycharmProjects\minCovTarget\venv\lib\site-packages\numba\np\linalg.py (817) raised from C:\Users\a_che\PycharmProjects\minCovTarget\venv\lib\site-packages\numba\core\errors.py:837 During: resolving callee type: Function(<function dot at 0x00000280CC542EF0>) During: typing of call at C:\Users\a_che\PycharmProjects\minCovTarget\venv\lib\site-packages\numba\np\linalg.py (460) File "venv\lib\site-packages\numba\np\linalg.py", line 460: def dot_impl(a, b): <source elided> out = np.empty((m, n), a.dtype) return np.dot(a, b, out) ^ During: lowering "$8call_function.3 = call $2load_deref.0(left, right, func=$2load_deref.0, args=[Var(left, linalg.py:582), Var(right, linalg.py:582)], kws=(), vararg=None, varkwarg=None, target=None)" at C:\Users\a_che\PycharmProjects\minCovTarget\venv\lib\site-packages\numba\np\linalg.py (582) raised from C:\Users\a_che\PycharmProjects\minCovTarget\venv\lib\site-packages\numba\core\typeinfer.py:1086 - Of which 2 did not match due to: Overload in function 'dot_3': File: numba\np\linalg.py: Line 784. With argument(s): '(array(int64, 2d, C), array(int64, 2d, C))': Rejected as the implementation raised a specific error: TypingError: missing a required argument: 'out' raised from C:\Users\a_che\PycharmProjects\minCovTarget\venv\lib\site-packages\numba\core\typing\templates.py:784 During: resolving callee type: Function(<function dot at 0x00000280CC542EF0>) During: typing of call at C:\Users\a_che\PycharmProjects\minCovTarget\tst4.py (164) File "tst4.py", line 164: def tst_dot(a, b): <source elided> return np.dot(a, b) ^
原因分析
从错误信息中的unsupported dtype for <BLAS function>()可以看出,核心问题是Numba对int64类型的二维数组矩阵乘法支持不完善。Numba底层依赖BLAS库实现矩阵运算,但部分BLAS版本对整数类型的矩阵乘法支持有限,导致无法找到匹配的实现逻辑。
解决方法
方法1:转换数组为浮点类型
将int64数组转为float32或float64,Numba对浮点型的矩阵乘法支持更完善:
import numpy as np import numba @numba.njit() def tst_dot(): a = np.array([[1, 0], [0, 1]], dtype=np.float64) b = np.array([[4, 1], [2, 2]], dtype=np.float64) return np.dot(a, b) print(tst_dot())
方法2:启用fastmath优化
使用@numba.njit(fastmath=True)装饰器,开启快速数学优化后,部分场景下能自动兼容整数矩阵乘法:
import numpy as np import numba @numba.njit(fastmath=True) def tst_dot(): a = np.array([[1, 0], [0, 1]]) b = np.array([[4, 1], [2, 2]]) return np.dot(a, b) print(tst_dot())
方法3:手动实现矩阵乘法(适合小矩阵)
如果矩阵规模不大,手动实现矩阵乘法可以绕过Numba的np.dot限制:
import numpy as np import numba @numba.njit() def manual_dot(a, b): m, n = a.shape p = b.shape[1] result = np.zeros((m, p), dtype=a.dtype) for i in range(m): for k in range(n): if a[i, k] != 0: for j in range(p): result[i, j] += a[i, k] * b[k, j] return result @numba.njit() def tst_dot(): a = np.array([[1, 0], [0, 1]]) b = np.array([[4, 1], [2, 2]]) return manual_dot(a, b) print(tst_dot())
验证结果
以上三种方法都能得到和无Numba时一致的输出:
- 方法1输出浮点型:
[[4. 1.] [2. 2.]] - 方法2、3输出整数型:
[[4 1] [2 2]]
内容的提问来源于stack exchange,提问作者Cherns
相关产品推荐
相关产品推荐

