如何在Numba中重新缓存全局变量?解决@njit编译后值无法修改问题
问题描述
尝试修改一个全局状态(比如OpenSimplex的随机种子),但该状态被@njit装饰的函数调用时,函数一旦编译,Numba就会固定全局值,导致后续修改无法生效。示例代码如下:
from numba import njit global_var = 3 @njit def func(): return global_var - 3 print(func()) # 输出0 global_var = 5 print(func()) # 仍输出0,不符合预期
尝试过的方法及问题
试过用闭包和numba.experimental.jitclass存储状态,但均未成功。jitclass示例代码:
from numba import njit, int32 from numba.experimental import jitclass spec = [ ('var', int32), ] @jitclass(spec) class State: def __init__(self, var): self.var = var def set(self, var): self.var = var state = State(1) @njit def get_state(): return state.var get_state() # 抛出错误
触发的错误信息:
Traceback (most recent call last):
File "D:\rnd\py\landslip\try-jit.py", line 20, in
get_state()
File "D:\prog\Python311\Lib\site-packages\numba\core\dispatcher.py", line 468, in _compile_for_args
error_rewrite(e, 'typing')
File "D:\prog\Python311\Lib\site-packages\numba\core\dispatcher.py", line 409, in error_rewrite
raise e.with_traceback(None)
numba.core.errors.NumbaNotImplementedError: Failed in nopython mode pipeline (step: native lowering)
<numba.core.base.OverloadSelector object at 0x00000206A85AC7D0>, (instance.jitclass.State#206a825e050var:int32,)
During: lowering "$4load_global.0 = global(state: <numba.experimental.jitclass.boxing.State object at 0x00000206A8597DC0>)" at D:\rnd\py\landslip\try-jit.py (18)
可行解决方案
方法1:将状态作为参数传递(最稳妥)
直接把需要修改的状态作为参数传入njit函数,完全规避全局变量依赖,符合Numba设计逻辑:
from numba import njit @njit def func(global_var): return global_var - 3 print(func(3)) # 输出0 print(func(5)) # 输出2,符合预期
方法2:用objmode临时绕过nopython模式(适合简单场景)
如果必须依赖全局变量,可在访问全局值的代码段用objmode包裹,让这部分回到Python解释器执行,读取最新全局值:
from numba import njit, objmode global_var = 3 @njit def func(): with objmode(var='int64'): var = global_var return var - 3 print(func()) # 输出0 global_var = 5 print(func()) # 输出2,符合预期
注意:objmode会带来性能损耗,适合对性能要求不高的场景。
方法3:使用Numba typed容器存储状态
用numba.typed.Dict或numba.typed.List存储状态,这类容器支持在njit函数中修改并保持状态,适合多函数共享场景:
from numba import njit from numba.typed import Dict from numba.core import types # 初始化typed字典存储状态 state = Dict.empty(types.unicode_type, types.int64) state['seed'] = 3 @njit def func(): return state['seed'] - 3 print(func()) # 输出0 state['seed'] = 5 print(func()) # 输出2,符合预期
方法4:修复jitclass的使用方式
之前的错误是因为直接在njit函数中访问全局jitclass实例,正确做法是将实例作为参数传入函数:
from numba import njit, int32 from numba.experimental import jitclass spec = [ ('var', int32), ] @jitclass(spec) class State: def __init__(self, var): self.var = var def set(self, var): self.var = var state = State(1) @njit def get_state(state_inst): return state_inst.var print(get_state(state)) # 输出1 state.set(5) print(get_state(state)) # 输出5,符合预期
内容的提问来源于stack exchange,提问作者Jonas Byström

