如何在Numba编译的函数及调用函数中实现结果缓存?
Numba环境下实现函数结果缓存的解决方案
原方案失败原因
- 先加
@njit再加@Memoize:此时bla是自定义Memoize类的实例,Numba编译caller时无法识别该自定义类型,触发类型推断错误。 - 调换装饰器顺序:
@njit仅支持装饰函数,无法装饰类实例,直接抛出类型错误。
可行方案
方案一:用jitclass封装缓存与计算逻辑
利用Numba实验性的jitclass将缓存字典和计算逻辑封装为Numba可识别的类,让缓存逻辑成为类方法,可在njit编译的调用函数中直接使用。
from numba import njit, types from numba.experimental import jitclass # 定义缓存字典的类型规范:键为(int, float)元组,值为float memo_dict_type = types.DictType(types.Tuple((types.int64, types.float64)), types.float64) spec = [("memo", memo_dict_type)] @jitclass(spec) class MemoizedCalculator: def __init__(self): self.memo = {} # 初始化空缓存 def compute(self, a: int, b: float): key = (a, b) if key in self.memo: return self.memo[key] # 原计算逻辑 result = a for i in range(1_000_000_000): result *= b self.memo[key] = result return result # 创建全局缓存计算器实例(Numba可识别) calculator = MemoizedCalculator() @njit def caller(x: int): s = 0.0 for j in range(x): # 调用缓存方法获取结果 s += calculator.compute(j % 5, float((j + 1) % 5)) return s if __name__ == "__main__": print(caller(30))
方案二:用staticvar在函数内部维护静态缓存
无需额外类,直接在njit编译的函数内部用静态变量保存缓存字典,Numba会在函数多次调用时复用该变量。
from numba import njit, types, staticvar @njit def bla(a: int, b: float): # 声明静态缓存字典并指定类型 memo = staticvar(types.DictType(types.Tuple((types.int64, types.float64)), types.float64)) # 首次调用时初始化字典 if memo is None: memo = {} key = (a, b) if key in memo: return memo[key] # 原计算逻辑 result = a for i in range(1_000_000_000): result *= b memo[key] = result return result @njit def caller(x: int): s = 0.0 for j in range(x): s += bla(j % 5, float((j + 1) % 5)) return s if __name__ == "__main__": print(caller(30))
注意事项
- 两种方案都必须明确缓存字典的键值类型,Numba在nopython模式下不支持动态类型字典。
- 因参数
b为float类型,调用时需确保传入值为float(如float((j+1)%5)),避免类型不匹配导致缓存失效。 jitclass属于Numba实验性功能,当前版本稳定可用,但未来可能存在API变动。
内容的提问来源于stack exchange,提问作者user107511
相关产品推荐
相关产品推荐

