在Numba jitclass中使用numpy.array初始化矩阵报错的原因及解决方案咨询
在Numba Jitclass中使用
np.array初始化二维矩阵的问题解决 我来帮你分析这个问题,并且给出更直接的解决方案——不用vstack也能直接用np.array完成初始化。
问题根源
你遇到的TypingError本质是Numba在nopython模式下,对混合类型的嵌套字面量数组初始化的类型推导出现了内部错误。看你的初始化代码:
self.n = np.array([[0.,1],[2,3],[4,5]])
这里的列表里同时存在浮点数(0.)和整数(1、2等),Numba在jitclass的构造函数中处理这种混合类型的嵌套结构时,无法正确推断最终数组的类型,从而触发了内部的类型推断失败。
直接修复方案
有两种简单的方法可以直接用np.array完成初始化,不需要借助vstack:
方案1:统一字面量的数值类型
把所有字面量都改成浮点数形式,让Numba能清晰推断出数组类型:
from numba.experimental import jitclass from numba import float64 import numpy as np spec = [('n', float64[:,:])] @jitclass(spec) class myclass(object): def __init__(self ): # 所有元素都用浮点数形式 self.n = np.array([[0., 1.], [2., 3.], [4., 5.]]) if __name__ == '__main__': pop = myclass() print(pop.n)
方案2:显式指定dtype参数
给np.array加上dtype=np.float64,强制指定数组的类型,跳过Numba的自动类型推断:
from numba.experimental import jitclass from numba import float64 import numpy as np spec = [('n', float64[:,:])] @jitclass(spec) class myclass(object): def __init__(self ): # 显式指定dtype为float64 self.n = np.array([[0.,1],[2,3],[4,5]], dtype=np.float64) if __name__ == '__main__': pop = myclass() print(pop.n)
为什么你的临时方案有效?
你用vstack的方法之所以能工作,是因为每个np.array([0., 1])这样的子数组都被显式创建为浮点数数组,vstack合并后的数组类型明确,Numba能正确识别,所以不会触发类型推断错误。但相比之下,上面两种直接用np.array的方案更简洁直观。
内容的提问来源于stack exchange,提问作者ymmx
相关产品推荐
相关产品推荐

