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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 08:39:55