Numba JitClass中二维数组预分配的报错问题及解决咨询
搞定Numba jitclass里二维数组的报错问题
我看了你的代码和报错信息,问题其实不是二维数组的声明方式,而是数组类型不匹配!np.zeros()默认生成的是float64类型的数组,但你在spec里给array和foo_matrix指定的是float32和int32,Numba的jitclass对类型一致性要求特别严,这就导致了类型推断失败。
具体修复方法很简单:
- 给
np.zeros()加上dtype参数,显式指定和spec里一致的类型; - 二维数组用元组
(value, value)定义形状是没问题的,不用改这个。
修正后的完整代码:
import numpy as np from numba import int32, float32 from numba.experimental import jitclass spec = [ ('value', int32), ('array', float32[:]), ('foo_matrix', int32[:,:]), ] @jitclass(spec) class Bag(object): def __init__(self, value): self.value = value # 显式指定float32,匹配spec里的array类型 self.array = np.zeros(value, dtype=np.float32) # 显式指定int32,匹配二维数组的spec定义 self.foo_matrix = np.zeros((value, value), dtype=np.int32) @property def size(self): return self.array.size def increment(self, val): for i in range(self.size): self.array[i] = val return self.array my_class = Bag(3) # 可以加个测试输出验证 print("Array:", my_class.array) print("Matrix:", my_class.foo_matrix)
再解释下报错的原因:
Numba的jitclass是静态类型的,spec里声明的每个属性类型必须和实际赋值的变量完全对应。np.zeros()默认会生成float64类型的数组,而你在spec里用了float32和int32,这种类型不匹配直接触发了Numba的类型检查错误。你之前改二维数组的形状写法(元组/列表)其实不是问题,Numba对这两种写法都支持,核心还是类型没对齐。
内容的提问来源于stack exchange,提问作者Siderius
相关产品推荐
相关产品推荐

