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
相关产品推荐
相关产品推荐

