You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.23 16:57:02