如何在Numba jitclass中正确构造函数指针数组
Numba构造函数指针数组的正确方案
原有代码的核心问题
- 漏导入
numpy包,直接触发NameError: name 'np' is not defined - Numpy没有对应函数指针的dtype,无法直接存储Numba函数指针,因此报
none cannot be represented as a Numpy dtype错误 - 函数签名(Signature)不是可索引的类型对象,不能直接加
[:]声明数组类型,因此报'Signature' object is not subscriptable错误 - jitclass方法直接作为指针存储会有循环类型依赖问题,需要先显式声明函数类型
目前Numba不支持将函数指针存入numpy数组,只能使用Numba专属的typed容器存储函数指针序列
正确实现代码
import numpy as np from numba import njit, void, deferred_type, float64 from numba.experimental import jitclass from numba.typed import List from numba.core.types import FunctionType, ListType # 先声明类的延迟类型,解决循环依赖问题 Test_type = deferred_type() # 定义函数指针类型:输入是Test类实例,无返回值 func_type = FunctionType(void(Test_type)) # 把方法逻辑拆成独立的njit函数,避免类型依赖问题 @njit(void(Test_type)) def x_func(self): self.a += 1 @njit(void(Test_type)) def y_func(self): self.a += 2 @jitclass(spec={ 'a': float64, 'ptrs': ListType(func_type) # 用numba typed List存储函数指针 }) class Test: def __init__(self): self.a = 0.0 # 初始化typed List存入两个函数指针 self.ptrs = List.empty_list(func_type) self.ptrs.append(x_func) self.ptrs.append(y_func) def increment(self, n): self.ptrs[n](self) # 绑定延迟类型的实际对应类 Test_type.define(Test.class_type.instance_type) # 运行测试 t = Test() print(t.a) # 输出:0 t.increment(0) print(t.a) # 输出:1 t.increment(1) print(t.a) # 输出:3
内容的提问来源于stack exchange,提问作者Mad Physicist
相关产品推荐
相关产品推荐

