Numba njit中jitclass对象列表的类型签名问题求解
Numba中njit函数接收jitclass对象列表的类型签名问题解决
单jitclass对象处理(正常运行)
仅传递单个jitclass对象时,代码可正常执行:
import numba as nb @nb.experimental.jitclass([('x', nb.float64)]) class test_jitclass(object): def __init__(self, x): self.x = x test_jitclass_type = test_jitclass.class_type.instance_type @nb.njit(nb.int64(test_jitclass_type)) def process_test_jitclass(tjc): print('processing: ', tjc.x) return 1 process_test_jitclass(test_jitclass(5.))
扩展至jitclass对象列表(出现错误)
尝试传递jitclass对象列表时,首次编写的代码如下:
import numba as nb @nb.experimental.jitclass([('x', nb.float64)]) class test_jitclass(object): def __init__(self, x): self.x = x test_jitclass_type = test_jitclass.class_type.instance_type test_jitclass_list_type = nb.types.ListType(test_jitclass_type) list_test_jitclass = [] for i in range(3): list_test_jitclass.append(test_jitclass(i)) @nb.njit(nb.int64(test_jitclass_list_type)) def process_list_test_jitclass(list_test_jitclass): for tj in list_test_jitclass: print('processing: ', tj.x) return 1 process_list_test_jitclass(list_test_jitclass)
执行后触发错误:
# Traceback (most recent call last): # File "<string>", line 1, in <module> # File "/home/____/.pyenv/versions/3.9.16/lib/python3.9/site-packages/numba/core/dispatcher.py", line 703, in _explain_matching_error # raise TypeError(msg) # TypeError: No matching definition for argument type(s) reflected list(instance.jitclass.test_jitclass#7f83a409d970<x:float64>)<iv=None>
尝试用nb.typed.List定义类型时,执行test_jitclass_list_type = nb.typed.List(test_jitclass_type)也会报错:
# Traceback (most recent call last): # File "<string>", line 7, in <module> # File "/home/____/.pyenv/versions/3.9.16/lib/python3.9/site-packages/numba/typed/typedlist.py", line 268, in __init__ # for i in args[0]: # File "/home/____/.pyenv/versions/3.9.16/lib/python3.9/site-packages/numba/core/types/abstract.py", line 185, in __getitem__ # ndim, layout = self._determine_array_spec(args) # File "/home/____/.pyenv/versions/3.9.16/lib/python3.9/site-packages/numba/core/types/abstract.py", line 210, in _determine_array_spec # raise KeyError(f"Can only index numba types with slices with no start or stop, got {args}.") # KeyError: 'Can only index numba types with slices with no start or stop, got 0.'
注:使用Python 3.9.16版本,无法更换。
可行解决方案
修改后的代码可正常运行:
import numba as nb @nb.experimental.jitclass([('x', nb.float64)]) class test_jitclass(object): def __init__(self, x): self.x = x test_jitclass_type = test_jitclass.class_type.instance_type list_test_jitclass = nb.typed.List([test_jitclass(i) for i in range(3)]) @nb.njit(nb.int64(nb.types.ListType(test_jitclass_type))) def process_list_test_jitclass(list_test_jitclass): for tjc in list_test_jitclass: print('processing: ', tjc.x) return 1 process_list_test_jitclass(list_test_jitclass)
关键修改点
- 对变量而非类型使用
nb.typed.List,创建Numba可识别的typed列表对象 - njit函数签名中依然使用
nb.types.ListType(test_jitclass_type)来声明参数类型
内容的提问来源于stack exchange,提问作者Squishium
相关产品推荐
相关产品推荐

