使用numba调用np.random.randint时报TypingError该如何解决
错误原因
numba的nopython模式不会完整适配所有numpy原生API的语法糖,你遇到的TypingError是因为低版本numba对np.random.randint的绑定存在限制:不支持省略low参数默认取0的写法,同时size参数只支持传入整数生成一维数组,不支持直接传入元组生成高维数组。
解决方案
以下三种方案都可以解决问题:
- 方案1:兼容低版本numba的写法,调整
np.random.randint的传参形式
from numba import jit import numpy as np @jit(nopython=True) def foo(): # 显式指定low=0,先生成一维数组再reshape为目标形状 a = np.random.randint(0, 16, size=9).reshape(3,3) return a foo()
- 方案2:使用numba官方推荐的随机数生成器接口,适配性和性能都更优
from numba import jit from numba.random import default_rng @jit(nopython=True) def foo(): rng = default_rng() # 直接支持size传入元组生成高维数组 a = rng.integers(0, 16, size=(3,3)) return a foo()
- 方案3:升级numba到0.57及以上版本,新版本已经完整适配
np.random.randint的省略low、size传元组的语法,你的原始代码可以直接正常运行。
内容的提问来源于stack exchange,提问作者Rice
相关产品推荐
相关产品推荐

