【发布时间】:2020-06-01 06:21:59
【问题描述】:
考虑以下函数来计算 3 n + 1 问题的给定输入的步数:
def num_steps(b, steps):
e = b
d = 0
while True:
if e == 1:
d += steps[e]
return d
if e % 2 == 0:
e //= 2
else:
e = 3*e + 1
d += 1
这里,steps 的存在是为了允许对结果进行记忆,但是为了这个问题,我们只注意只要steps[1] == 0,它应该没有效果,因为在这种情况下,效果d += steps[e] 是在d 上加0。事实上,下面的例子给出了预期的结果:
import numpy as np
steps = np.array([0, 0, 0, 0])
print(num_steps(3, steps)) # Prints 7
但是,如果我们使用 numba.jit(或 njit)对方法进行 JIT 编译,我们将不再得到正确的结果:
import numpy as np
from numba import jit
steps = np.array([0, 0, 0, 0])
print(jit(num_steps)(3, steps)) # Prints 0
如果我们在编译方法之前删除看似冗余的d += steps[e],我们确实会得到正确的结果。我们甚至可以在d += steps[e] 之前放入print(steps[e]) 并看到值为0。我还可以将d += 1 移动到循环的顶部(并初始化d = -1)以获得同样有效的东西在 Numba 案例中。
这发生在 Python 3.8 上的 Numba 0.48.0 (llvmlite 0.31.0)(通过标准 conda 渠道提供的最新版本)。
【问题讨论】:
-
@MrFuppes:谢谢,如果我使用 conda 将 Python 版本固定到 3.7(将 Numba 降级到 0.47.0),我仍然会遇到同样的问题。
-
抱歉,在我写初始评论之前没有检查我的 numba 版本 - 也可以使用 Python 3.7(numba 0.48,llvmlite 0.31)重现无效结果! 但是,如果我在 Python 3.7 上切换到 numba 0.46 / llvmlite 0.30,njitted 代码可以正常工作。
-
@MrFuppes:谢谢,我可以确认降级到 numba 0.46 / llvmlite 0.30 可以解决问题。
标签: python numpy jit numba collatz