【问题标题】:generator argument in NumbaNumba 中的生成器参数
【发布时间】:2017-06-11 04:30:52
【问题描述】:

对此问题的跟进:function types in numba

我正在编写一个需要将生成器作为其参数之一的函数。在这里粘贴太复杂了,所以考虑这个玩具示例:

def take_and_sum(gen):
    @numba.jit(nopython=False)
    def inner(n):
        s = 0
        for _ in range(n):
            s += next(gen)
        return s
    return inner

它返回生成器的第一个 n 元素的总和。用法示例:

@numba.njit()
def odd_numbers():
    n = 1
    while True:
        yield n
        n += 2

take_and_sum(odd_numbers())(3) # prints 9

这是 curried,因为我想用 nopython=True 编译,然后我不能将 genpyobject)作为参数传递。不幸的是,nopython=True 出现错误:

TypingError: Failed at nopython (nopython frontend)
Untyped global name 'gen'

即使我 nopython 编译了我的生成器。

真正令人困惑的是,对输入进行硬编码是可行的:

def take_and_sum():
    @numba.njit()
    def inner(n):
        gen = odd_numbers()
        s = 0.0
        for _ in range(n):
            s += next(gen)
        return s
    return inner

take_and_sum()(3)

我也尝试将我的生成器变成一个类:

@numba.jitclass({'n': numba.uint})
class Odd:
    def __init__(self):
        self.n = 1
    def next(self):
        n = self.n
        self.n += 2
        return n

同样,这在对象模式下有效,但在 nopython 模式下,我得到了无法搜索的结果:

LoweringError: Failed at nopython (nopython mode backend)
Internal error:
NotImplementedError: instance.jitclass.Odd#4aa9758<n:uint64> as constant unsupported

【问题讨论】:

    标签: python generator numba


    【解决方案1】:

    我实际上无法解决您的问题,因为据我所知根本不可能。我只是强调一些方面(适用于numba 0.30):

    不能创建一个 numba-jitclass 生成器:

    import numba
    
    @numba.jitclass({'n': numba.uint})
    class Odd:
        def __init__(self):
            self.n = 1
    
        def __iter__(self):
            return self
    
        def __next__(self):
            n = self.n
            self.n += 2
            return n
    

    试试吧:

    >>> next(Odd())
    TypeError: 'Odd' object is not an iterator
    

    当您删除 numba.jitclass 时,它会起作用:

    >>> next(Odd())
    1
    

    您使用硬编码生成器的示例不等效。您最初的尝试创建了一个生成器对象,并将其传递给一个 numba 函数并修改了生成器。您会期望它更新生成器的状态

    >>> t = odd_numbers()
    >>> take_and_sum(t)(3)
    9
    >>> next(t)   # State has been updated, unfortunatly that requires nopython=False!
    7
    

    但这对于 numba 来说是不可能的(还)。

    第二个例子不同,因为你每次调用函数时都会创建生成器,所以你的函数之外没有需要更新的状态:

    >>> take_and_sum()(3) # using your hardcoded version
    9.0
    >>> take_and_sum()(3) # no updated state so this returns the same:
    9.0
    

    绝对可以更改它,但没有使用任意函数的选项:

    @numba.jitclass({'n': numba.uint})
    class Odd:
        def __init__(self):
            self.n = 1
    
        def calculate(self, n):
            s = 0.0
            for _ in range(n):
                s += self.n
                self.n += 2
            return s
    
    >>> x = Odd()
    >>> x.calculate(3)
    9.0
    >>> x.calculate(3)
    27.0
    

    我知道这不是你想要的,但至少它是可行的 :-)

    【讨论】:

    • 是的,我想说我大多数 numba 问题的根本原因是无法将非局部变量捕获为可写
    猜你喜欢
    • 2011-06-05
    • 2020-08-23
    • 2013-04-18
    • 1970-01-01
    • 2015-12-07
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多