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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 07:06:03