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

Numba jitclass不支持元组键字典?@jit可运行jitclass失效

问题说明

使用@njit装饰函数时,以元组为键的字典能正常运行,但改用jitclass实现时,Dict.empty相关逻辑无法工作。疑问:是不是jitclass不支持元组作为字典键,而@jit支持该特性?

原问题代码

from numba import types, njit
from numba.experimental import jitclass
from numba.typed import Dict
from numba import int64, float64

spec = [
    #key : time, s
    ('values', types.DictType(types.Tuple((int64, types.unicode_type)), float64))
]

@jitclass(spec)
class Data(object):
    def __init__(self):
        self.values = Dict.empty(
            key_type=types.Tuple((int64, types.unicode_type)),
            value_type=float64,
        )

    def add_value(self, time: int, s: str, value: float):
        self.values[(time, s)] = value

    def get_value(self, time: int, s: str):
        return self.values.get((time, s), None)

    def iterate_values(self):
        for (time, s), value in self.values.items():
            print(f"At time {time}, s {s} had a value of {value}")

# Instantiate the class
data = Data()

# Add some sample data
data.add_value(20230530, "s", 150.45)

print(data.get_value(20230530, "s"))  # Outputs: 150.45

# Iterate over all values
data.iterate_values()
解答

没错,你猜的完全正确——numba的jitclass目前确实不支持以元组作为键的字典,但njit(@jit的no-python模式)是支持这个特性的。

这是因为jitclass的类型校验、内存布局机制和普通njit函数不同,对复合类型(比如元组)作为字典键的场景支持存在兼容性缺陷,属于numba已知的功能限制。

可行解决方案

最直接的替代方案是把元组键转换成唯一字符串键,比如将时间和字符串拼接,这样就能在jitclass里正常使用字典了。修改后的代码如下:

from numba import types, njit
from numba.experimental import jitclass
from numba.typed import Dict
from numba import int64, float64

# 修改spec:字典键改为字符串类型
spec = [
    ('values', types.DictType(types.unicode_type, float64))
]

@jitclass(spec)
class Data(object):
    def __init__(self):
        self.values = Dict.empty(
            key_type=types.unicode_type,
            value_type=float64,
        )

    def add_value(self, time: int, s: str, value: float):
        # 将元组转为拼接字符串作为键
        key = f"{time}_{s}"
        self.values[key] = value

    def get_value(self, time: int, s: str):
        key = f"{time}_{s}"
        return self.values.get(key, None)

    def iterate_values(self):
        for key, value in self.values.items():
            # 拆分字符串还原原数据
            time_str, s = key.split('_')
            time = int(time_str)
            print(f"At time {time}, s {s} had a value of {value}")

# 测试代码
data = Data()
data.add_value(20230530, "s", 150.45)
print(data.get_value(20230530, "s"))  # 输出: 150.45
data.iterate_values()

如果不想用字符串拼接,也可以尝试用numba的StructRef自定义结构体类型作为键,但实现复杂度会高很多,字符串拼接是最省心的折中方案。

内容的提问来源于stack exchange,提问作者tompal18

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 13:27:33