向Numba函数传入数组时函数签名不匹配的问题排查
解决Numba函数二维数组参数签名不匹配问题
错误原因
你当前的函数签名float64(float64[:])定义了接受一维浮点数组,返回单个浮点标量,但实际传入的是二维数组array(float64, 2d, C),维度不匹配导致类型错误。
正确解决方案
方案1:直接适配二维数组输入,返回一维结果
修改函数签名为float64[:](float64[:,:]),表示接受二维浮点数组,返回一维浮点数组。同时优化函数内部逻辑(原计算式等价于(a0+a1)²):
import numpy as np from numba import njit, float64 A = [( 0.0182286178413157, -1.2904019395416308), ( 0.5228683581098151, 0.2323207738837293), (-0.6056770113345468, 1.5990251249135883), (-0.7557841434090988, 1.4641641762952791), ( 0.9882455737412416, -1.1838797980930709), (-1.2168205368640061, 1.5178083863904257), (-0.5566781056044838, 0.2160324328998916), ( 0.0671405605855369, -0.4246242749812621), ( 0.4806167193998933, 1.0521631181457611), ( 0.0563547059786364, -0.8223422191733811)] A = np.array(A) @njit(float64[:](float64[:,:])) def distance(a): result = np.empty(a.shape[0], dtype=np.float64) for i in range(a.shape[0]): result[i] = (a[i,0] + a[i,1]) ** 2 return result print(distance(A))
方案2:保留一维数组处理函数,封装二维数组调用
如果希望保留原函数处理单行数据的逻辑,可以额外封装一个适配二维输入的函数:
import numpy as np from numba import njit, float64 A = [( 0.0182286178413157, -1.2904019395416308), ( 0.5228683581098151, 0.2323207738837293), (-0.6056770113345468, 1.5990251249135883), (-0.7557841434090988, 1.4641641762952791), ( 0.9882455737412416, -1.1838797980930709), (-1.2168205368640061, 1.5178083863904257), (-0.5566781056044838, 0.2160324328998916), ( 0.0671405605855369, -0.4246242749812621), ( 0.4806167193998933, 1.0521631181457611), ( 0.0563547059786364, -0.8223422191733811)] A = np.array(A) @njit(float64(float64[:])) def distance_row(a): return (a[0] + a[1]) ** 2 @njit(float64[:](float64[:,:])) def distance_matrix(mat): res = np.empty(mat.shape[0], dtype=np.float64) for i in range(mat.shape[0]): res[i] = distance_row(mat[i]) return res print(distance_matrix(A))
签名规则说明
Numba的数组签名格式:
float64[:]:一维浮点数组float64[:,:]:二维浮点数组float64:单个浮点标量
签名结构为返回类型(参数类型),只要确保参数类型和实际传入的数组维度、类型一致即可。
内容的提问来源于stack exchange,提问作者Sergio Vivas Pleite
相关产品推荐
相关产品推荐

