使用Numba JIT(nopython=True)时numpy linspace函数报错,求解决方法
解决Numba nopython模式下linspace生成整数序列报错的问题
为啥linspace会报错?
Numba的nopython=True模式对np.linspace的整数参数支持有限——毕竟linspace本身是为生成浮点序列设计的,你强行把所有参数都转成np.int_,会打乱Numba的类型推断逻辑,直接触发一堆错误。
最简单的可行方案
要生成从0到Nt-1的整数序列,直接用np.arange就好,这货专门用来生成等差整数序列,Numba对它的支持非常完善,代码还更简洁:
from numba import jit import numpy as np @jit(nopython=True) def func(Nt): time = np.arange(Nt, dtype=np.int_) return time Nt = 10 print(func(Nt)) # 输出: [0 1 2 3 4 5 6 7 8 9]
非要用linspace的话(不推荐)
如果执念要用linspace,别手动转换参数类型,先让它生成浮点序列再转成整数,这样Numba能正常处理:
@jit(nopython=True) def func_linspace(Nt): time = np.linspace(0, Nt-1, Nt).astype(np.int_) return time print(func_linspace(10)) # 同样能得到目标整数序列
内容的提问来源于stack exchange,提问作者Aud
相关产品推荐
相关产品推荐

