【问题标题】:Numba jitting changes result when adding certain kind of 0 to local variable将某种 0 添加到局部变量时,Numba jitting 会改变结果
【发布时间】: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


【解决方案1】:

对我来说,这看起来像是一个错误,带有steps[e] 的就地增量。如果你设置parallel=True 那就是 Numba 崩溃的地方。你可以在 Numba github repo 上创建一个问题,也许开发人员可以解释一下。

如果我重写函数以避免最终的就地增量,它对我有用:

@numba.njit
def numb_steps(b, steps):

    e = b    
    d = 0

    while True:

        if e == 1:
            return d + steps[e]

        if e % 2 == 0:
            e //= 2
        else:
            e = 3*e + 1

        d += 1

与:

python                    3.7.6
numba                     0.47.0

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 2018-11-13
    • 2020-12-06
    • 2022-12-06
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多