You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Python Numba njit模式加速动态规划代码TypingError报错求助

问题描述

我是Python初学者,计划使用Python开展数值实验,实验过程中需要精确求解大量动态规划问题,因此代码运行效率的优化至关重要。目前我的代码搭配Numba的@jit装饰器可正常运行,但我希望进一步使用@njit装饰器提升运行性能。我已经尝试对for循环内的运算做向量化处理以提升效率,但@jit模式可正常运行的代码切换为@njit模式后持续抛出错误。精确求解动态规划属于计算密集型任务,我非常希望能够通过@njit进一步提升性能,求可适配@njit模式的代码修改方案。

原实现代码
import numba as nb
import numpy as np

#DP computation
@nb.njit
def dp(beta,cost,wcost,decisions,number_of_stages,states):
    tbeta=1-beta
    odcost=min((cost[-max(decisions):]+wcost)/beta[-max(decisions):])
    terminal=(max(states)-states)*odcost
    L=max(states)
    D=number_of_stages
    value=np.zeros((D+1,L+1))
    choice=np.zeros((D+1,L)).astype(np.int64)
    value[-1]=terminal
    for s in range(D-1,L-2,-1):
        intmatrix=cost[:, None]+np.outer(beta,value[s+1][1:L+1])+np.outer(tbeta,value[s+1][0:L])
        choice[s]=intmatrix.T.argmin(axis=1)
        value[s][0:L]=intmatrix[choice[s],np.arange(intmatrix.shape[1])]
    
    for s in range(L-2,-1,-1):
        intmatrix=cost[:, None]+np.outer(beta,value[s+1][1:s+2])+np.outer(tbeta,value[s+1][0:s+1])
        choice[s][0:s+1]=intmatrix.T.argmin(axis=1)
        value[s][0:s+1]=intmatrix[choice[s][0:s+1],np.arange(intmatrix.shape[1])]
        
    return value, choice


#initialization
decisions=np.arange(100)
number_of_stages=200
states=np.arange(101)

np.random.seed(2021)
beta=np.append(0,np.random.uniform(0,1,max(decisions)))
wcost=np.random.uniform(0,1)
cost=np.square(beta)

value, choice=dp(beta,cost,wcost,decisions,number_of_stages,states)
报错信息
TypingError: No implementation of function Function(<built-in function getitem>) found for signature:
 
getitem(array(float64, 1d, C), Tuple(slice<a:b>, none))
 
There are 22 candidate implementations:
      - Of which 20 did not match due to:
      Overload of function 'getitem': File: <numerous>: Line N/A.
        With argument(s): '(array(float64, 1d, C), Tuple(slice<a:b>, none))':
       No match.
      - Of which 2 did not match due to:
      Overload in function 'GetItemBuffer.generic': File: numba\core\typing\arraydecl.py: Line 162.
        With argument(s): '(array(float64, 1d, C), Tuple(slice<a:b>, none))':
       Rejected as the implementation raised a specific error:
         TypeError: unsupported array index type none in Tuple(slice<a:b>, none)
  raised from C:\ProgramData\Anaconda3\lib\site-packages\numba\core\typing\arraydecl.py:68
问题原因与修改方案

报错核心原因是Numba的@njit(nopython模式)不支持使用None作为数组索引来扩展维度,代码中cost[:, None]的写法就是通过None给一维数组新增第二个维度,这个写法在普通NumPy、@jit的object模式下可以正常运行,但在nopython模式下无法被识别。

需要做两处修改即可适配@njit:

  • 将所有[:, None]的维度扩展写法,替换为@njit支持的reshape操作,cost.reshape(-1, 1)和cost[:, None]效果完全一致
  • 数组初始化时直接指定dtype,替代先创建数组再调用astype转类型的写法,执行效率更高

修改后可正常运行的代码如下:

import numba as nb
import numpy as np

#DP computation
@nb.njit
def dp(beta,cost,wcost,decisions,number_of_stages,states):
    tbeta=1-beta
    max_dec = max(decisions)
    odcost=min((cost[-max_dec:]+wcost)/beta[-max_dec:])
    terminal=(max(states)-states)*odcost
    L=max(states)
    D=number_of_stages
    value=np.zeros((D+1,L+1))
    # 初始化时直接指定dtype,避免后续类型转换
    choice=np.zeros((D+1,L), dtype=np.int64)
    value[-1]=terminal
    # 提前处理列向量,避免循环内重复做维度转换
    cost_col = cost.reshape(-1, 1)
    for s in range(D-1,L-2,-1):
        intmatrix=cost_col + np.outer(beta,value[s+1][1:L+1])+np.outer(tbeta,value[s+1][0:L])
        choice[s]=intmatrix.T.argmin(axis=1)
        value[s][0:L]=intmatrix[choice[s],np.arange(intmatrix.shape[1])]
    
    for s in range(L-2,-1,-1):
        intmatrix=cost_col + np.outer(beta,value[s+1][1:s+2])+np.outer(tbeta,value[s+1][0:s+1])
        choice[s][0:s+1]=intmatrix.T.argmin(axis=1)
        value[s][0:s+1]=intmatrix[choice[s][0:s+1],np.arange(intmatrix.shape[1])]
        
    return value, choice


#initialization
decisions=np.arange(100)
number_of_stages=200
states=np.arange(101)

np.random.seed(2021)
beta=np.append(0,np.random.uniform(0,1,max(decisions)))
wcost=np.random.uniform(0,1)
cost=np.square(beta)

value, choice=dp(beta,cost,wcost,decisions,number_of_stages,states)

修改后的代码可以正常在@njit模式下运行,相比原@jit模式有明显的性能提升。


内容的提问来源于stack exchange,提问作者Vergil

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.03 02:01:24