【发布时间】:2020-12-24 01:06:11
【问题描述】:
我正在尝试在 @guvectorize 中调用 @guvectorize,但出现错误提示:
Untyped global name 'regNL_nb': cannot determine Numba type of <class 'numpy.ufunc'>
File "di.py", line 12:
def H2Delay_nb(S1, S2, R2):
H1 = regNL_nb(S1, S2)
^
这是一个 MRE:
import numpy as np
from numba import guvectorize, float64, int64, njit, cuda, jit
@guvectorize(["float64[:], float64[:], float64[:]"], '(n),(n)->(n)')
def regNL_nb(S1, S2, h2):
for i in range(len(S1)):
h2[i] = S1[i] + S2[i]
@guvectorize(["float64[:], float64[:], float64[:]"], '(n),(n)->(n)',nopython=True)
def H2Delay_nb(S1, S2, R2):
H1 = regNL_nb(S1, S2)
H2 = regNL_nb(S1, S2,)
for i in range(len(S1)):
R2[i] = H1[i] + H2[i]
S1 = np.array([1,2,3,4,5,6,7,8,9])
S2 = np.array([1,2,3,4,5,6,7,8,9])
H2 = H2Delay_nb(S1, S2)
print(H2)
我不知道如何告诉 numba 函数 regNL_nb 是一个 guvectorize 函数。
【问题讨论】:
-
预期的输出应该是
[ 4. 8. 12. 16. 20. 24. 28. 32. 36.],对吗? -
是的。每个数字只有四次
-
我按照原来的方式执行了您上面的脚本,除了将参数
nopython设置为False,因此有代码会退回到对象模式的风险——除了警告之外它运行良好。>>> print(H2)给出输出:[ 4. 8. 12. 16. 20. 24. 28. 32. 36.]。 -
是的,它是这样工作的,但是速度较慢。重点是保持nopython=True模式,这样代码运行的更快。
-
我知道;但是,如果您停用对象模式,则并非所有值都将作为 Python 对象处理。 (见numba.pydata.org/numba-doc/latest/…)