使用Numba @jit装饰器却触发nopython模式错误的求助
我在用值函数迭代求解多状态复杂动态规划问题,想通过Numba的@jit装饰器加速代码(最终要并行化循环)。但使用@jit时没开nopython模式,却收到了nopython相关的TypingError。
代码示例
@jit def vfi(cm, λ=1): vf_new = np.zeros_like(cm.vf) k_prime = np.zeros_like(cm.k_opt) for k_i in range(cm.k_grid_size): cm.set_k(cm.kgrid[k_i]) for b_i in range(cm.cb_grid_size): cm.set_b(cm.cb_mesh[0][b_i]) for s_i in range(2): cm.set_s(s_i) b_prime = cm.b_prime(cm.kgrid) vf_interp = RegularGridInterpolator((cm.kgrid, cm.cb_mesh[0],cm.sgrid), cm.vf) objective = cm.F - cm.T(cm.kgrid) - cm.C(cm.kgrid) + cm.β*cm.p*np.array([vf_interp(x) for x in zip(cm.kgrid,b_prime,np.zeros_like(b_prime))]) + cm.β*(1-cm.p)*np.array([vf_interp(x) for x in zip(cm.kgrid,b_prime,np.ones_like(b_prime))]) vf_new[k_i,b_i,s_i] = np.max(objective) k_prime[k_i,b_i,s_i] = np.argmax(objective) error = np.max(np.abs(cm.vf - vf_new)) cm.vf = cm.vf + λ*(vf_new-cm.vf) cm.k_opt = k_prime return error qe.util.tic() cm = CarrybacksModel() error = 10000000 itern=0 tol = 1e-5 while error>tol: error = vfi(cm) itern+=1 print(f"Iteration number {itern}, error = {error}.") print(f"Completed in {itern} iterations.") qe.util.toc()
报错信息
> --------------------------------------------------------------------------- TypingError Traceback (most recent call last) Cell In[57], line 7 5 tol = 1e-5 6 while error>tol: ----> 7 error = vfi(cm) 8 itern+=1 9 print(f"Iteration number {itern}, error = {error}.") File ~\miniconda3\Lib\site-packages\numba\core\dispatcher.py:468, in _DispatcherBase._compile_for_args(self, *args, **kws) 464 msg = (f"{str(e).rstrip()} \n\nThis error may have been caused " 465 f"by the following argument(s):\n{args_str}\n") 466 e.patch_message(msg) ---> 468 error_rewrite(e, 'typing') 469 except errors.UnsupportedError as e: 470 # Something unsupported is present in the user code, add help info 471 error_rewrite(e, 'unsupported_error') File ~\miniconda3\Lib\site-packages\numba\core\dispatcher.py:409, in _DispatcherBase._compile_for_args.<locals>.error_rewrite(e, issue_type) 407 raise e 408 else: ---> 409 raise e.with_traceback(None) TypingError: Failed in nopython mode pipeline (step: nopython frontend) Untyped global name 'RegularGridInterpolator': Cannot determine Numba type of <class 'type'> File "..\..\..\..\AppData\Local\Temp\ipykernel_16688\910653365.py", line 12: <source missing, REPL/exec in use?> This error may have been caused by the following argument(s): - argument 0: Cannot determine Numba type of <class '__main__.CarrybacksModel'>
我知道没法用njit/nopython模式,因为自定义了CarrybacksModel类,且scipy的RegularGridInterpolator不兼容。但我以为用@jit而非@njit不会触发nopython模式,为什么还报错?我试过自己实现用@jit装饰的插值函数,还是报错,求解决建议。
为什么@jit会触发nopython模式报错?
Numba的@jit装饰器默认行为是先尝试编译为nopython模式,只有当nopython模式编译失败时,才会回退到object模式。但如果在nopython模式的类型检查阶段就遇到无法处理的对象(比如自定义类、未支持的第三方库函数),会直接抛出TypingError,不会自动回退到object模式。
解决建议
1. 显式指定object模式
直接在@jit中添加nopython=False参数,强制Numba使用object模式,跳过nopython模式的类型检查:
@jit(nopython=False) def vfi(cm, λ=1): # 原函数内容不变
注意:object模式的加速效果远不如nopython模式,因为它仍会调用Python对象的操作,但至少能让代码正常运行。
2. 拆分代码,分离可nopython化的部分
把函数中纯数值计算、无自定义类/不兼容函数的部分单独提取出来,用@njit装饰,而自定义类的操作、插值等留在主函数中。这样既利用nopython模式的高效加速,又避开不兼容的部分。
比如,把目标函数的计算、max/argmax操作拆成单独的njit函数:
@njit def compute_objective(F_val, T_vals, C_vals, β, p, vf_interp_vals): term1 = β * p * vf_interp_vals[:, 0] term2 = β * (1 - p) * vf_interp_vals[:, 1] objective = F_val - T_vals - C_vals + term1 + term2 max_val = np.max(objective) argmax_idx = np.argmax(objective) return max_val, argmax_idx # 主函数用object模式 @jit(nopython=False) def vfi(cm, λ=1): vf_new = np.zeros_like(cm.vf) k_prime = np.zeros_like(cm.k_opt) # 提前创建插值器,避免循环内重复创建(大幅提升效率) vf_interp = RegularGridInterpolator((cm.kgrid, cm.cb_mesh[0], cm.sgrid), cm.vf) for k_i in range(cm.k_grid_size): cm.set_k(cm.kgrid[k_i]) F_val = cm.F T_vals = cm.T(cm.kgrid) C_vals = cm.C(cm.kgrid) for b_i in range(cm.cb_grid_size): cm.set_b(cm.cb_mesh[0][b_i]) for s_i in range(2): cm.set_s(s_i) b_prime = cm.b_prime(cm.kgrid) # 批量计算插值点,替代列表推导式 interp_points_0 = np.column_stack((cm.kgrid, b_prime, np.zeros_like(b_prime))) vf_interp_0 = vf_interp(interp_points_0) interp_points_1 = np.column_stack((cm.kgrid, b_prime, np.ones_like(b_prime))) vf_interp_1 = vf_interp(interp_points_1) vf_interp_vals = np.column_stack((vf_interp_0, vf_interp_1)) # 调用njit优化的计算函数 max_val, argmax_idx = compute_objective(F_val, T_vals, C_vals, cm.β, cm.p, vf_interp_vals) vf_new[k_i,b_i,s_i] = max_val k_prime[k_i,b_i,s_i] = argmax_idx error = np.max(np.abs(cm.vf - vf_new)) cm.vf = cm.vf + λ*(vf_new-cm.vf) cm.k_opt = k_prime return error
3. 替换不兼容的插值器
- 自行实现Numba兼容的插值函数:比如写一个线性插值函数并用
@njit装饰,完全避开scipy的RegularGridInterpolator。 - 使用
numba-scipy扩展:如果安装了该库,它支持部分scipy函数的nopython模式编译,包括RegularGridInterpolator(需注意版本兼容性)。
4. 避免循环内重复创建插值器
你的代码在三重循环内每次都创建RegularGridInterpolator,这是极大的性能浪费。不管用不用Numba,都应该把插值器的创建移到循环外面,只初始化一次,这不仅能减少开销,还能降低Numba处理的复杂度。
内容的提问来源于stack exchange,提问作者linkspan

