如何为numba jitclass的部分方法保留原生Python实现?
解决numba jitclass保留部分原生Python方法的方案
报错原因
numba的jitclass装饰器默认会编译类内定义的所有方法,且默认运行在nopython模式下,无法推断pandas.DataFrame等没有对应numba内置类型的Python对象类型,因此调用包含这类对象的方法会抛出类型错误。
解决方案
方案1:包装类代理(推荐,改造成本低、兼容性强)
将需要编译的逻辑拆分到内部私有jitclass中,对外暴露普通Python类作为代理,完全保留原有API,不需要修改任何外部调用代码:
from numba import typeof from numba.experimental import jitclass import numpy as np import pandas as pd # 内部私有jitclass,仅存放需要加速的属性和方法 @jitclass([('_some_attribute', typeof(1))]) class _MyClassJIT: _some_attribute: int def __init__(self, some_attribute): self._some_attribute = some_attribute def do_stuff_fast(self, x: np.array): # 用于内部循环的高速方法,由numba编译 return x[self._some_attribute] # 对外暴露的普通Python类,API与原实现完全一致 class MyClass: def __init__(self, some_attribute): self._jit_instance = _MyClassJIT(some_attribute) self._some_attribute = some_attribute # 代理加速方法,外部调用无感知 def do_stuff_fast(self, x: np.array): return self._jit_instance.do_stuff_fast(x) # 原生Python方法,可任意使用第三方库,无需编译 def do_stuff_slow(self, df: pd.DataFrame): return df["toto"].values[self._some_attribute] # 原有调用代码无需任何修改 instance = MyClass(5) print(instance.do_stuff_fast(np.ones(10)) ) df = pd.DataFrame( {'toto' :np.zeros(10) } ) print( instance.do_stuff_slow(df) )
适用场景:非编译方法数量多、逻辑复杂、依赖大量不兼容numba的第三方库的场景。
方案2:objmode上下文包裹
如果不想拆分类结构,可在需要运行原生Python逻辑的方法中用objmode上下文包裹对应代码,手动指定返回值的numba类型:
from numba import typeof, objmode from numba.experimental import jitclass import numpy as np import pandas as pd @jitclass([('_some_attribute', typeof(1))]) class MyClass: _some_attribute: int def __init__(self, some_attribute): self._some_attribute = some_attribute def do_stuff_fast(self, x: np.array): return x[self._some_attribute] def do_stuff_slow(self, df: pd.DataFrame): # 用objmode包裹原生Python逻辑,手动指定返回值类型 with objmode(ret='float64[:]'): ret = df["toto"].values[self._some_attribute] return ret
适用场景:仅存在少量简单原生逻辑的场景,缺点是需要提前明确返回值的numba类型,复杂逻辑可能依然会出现类型推断错误。
内容的提问来源于stack exchange,提问作者user3780669
相关产品推荐
相关产品推荐

