【问题标题】:How make a python class jitclass compatible when it contains itself jitclass classes?当 python 类 jitclass 包含自己的 jitclass 类时,如何使它兼容?
【发布时间】:2019-12-10 23:08:59
【问题描述】:

我正在尝试创建一个类,它可能是 jitclass 的一部分,但有一些属性是它们自己的 jitclass 对象。

例如,如果我有两个带有装饰器 @jitclass 的类,我希望将它们实例化为第三个类 (combined)。

import numpy as np
from numba import jitclass
from numba import boolean, int32, float64,uint8

spec = [
    ('type' ,int32),
    ('val' ,float64[:]),
    ('result',float64)]

@jitclass(spec)
class First:
    def __init__(self):
        self.type = 1
        self.val = np.ones(100)
        self.result = 0.
    def sum(self):
        self.result = np.sum(self.val)

@jitclass(spec)
class Second:
    def __init__(self):
        self.type = 2
        self.val = np.ones(100)
        self.result = 0.
    def sum(self):
        self.result = np.sum(self.val)



@jitclass(spec)
class Combined:
    def __init__(self):
        self.List = []
        for i in range(10):
            self.List.append(First())
            self.List.append(Second())

    def sum(self):
        for i, c in enumerate(self.List):
            c.sum()
    def getresult(self):
        result = []
        for i, c in enumerate(self.List):
            result.append(c.result)
        return result


C = Combined()
C.sum()
result = C.getresult()
print(result)

在该示例中,我收到一个错误,因为 numba 无法确定 self.List 的类型,它是两个 jitclass 的组合。

如何使Combined 类与jitclass 兼容?

更新

它尝试了我在其他地方找到的东西:

import numpy as np
from numba import jitclass, deferred_type
from numba import boolean, int32, float64,uint8
from numba.typed import List

spec = [
    ('type' ,int32),
    ('val' ,float64[:]),
    ('result',float64)]

@jitclass(spec)
class First:
    def __init__(self):
        self.type = 1
        self.val = np.ones(100)
        self.result = 0.
    def sum(self):
        self.result = np.sum(self.val)
 

 
spec1 = [('ListA',  List(First.class_type.instance_type, reflected=True))]

@jitclass(spec1)
class Combined:
    def __init__(self):
        self.ListA = [First(),First()] 

    def sum(self):
        for i, c in enumerate(self.ListA):
            c.sum()
    def getresult(self):
        result = []
        for i, c in enumerate(self.ListA):
            result.append(c.result)
        return result


C = Combined()
C.sum()
result = C.getresult()
print(result)

但我得到了这个错误

List(First.class_type.instance_type)
TypeError: __init__() takes 1 positional argument but 2 were given

【问题讨论】:

  • 你检查答案here了吗?这不是您需要的 100%,但它可能会有所帮助 - 它提供了一些示例。
  • 如果我只实例化一次类 First,这个解决方案就可以工作,但我想要一个列表。我更新了我的帖子。

标签: python class jit numba


【解决方案1】:

TL;DR:

  • 您可以在jitclass 中引用其他jitclasses,即使您有这些列表。您只需要更正命名空间numba.typed -> numba.types
  • 目前(从 numba 0.46 开始)在 jitclasses 或 no-python numba.jit 函数中不可能有异构列表。因此,您不能将FirstSecond 的两个实例都附加到同一个列表中。

解决numba.typed.List 异常

您的更新几乎是正确的。您需要使用numba.types.List 而不是numba.typed.List。区别有点微妙,但numba.types 包含签名类型,而numba.typed 命名空间包含可以在代码中实例化和使用的类。

所以如果你使用它会起作用:

spec1 = [('ListA',  nb.types.List(First.class_type.instance_type, reflected=True))]

更改此代码:

import numpy as np
import numba as nb

spec = [
    ('type', nb.int32),
    ('val', nb.float64[:]),
    ('result', nb.float64)
]

@nb.jitclass(spec)
class First:
    def __init__(self):
        self.type = 1
        self.val = np.ones(100)
        self.result = 0.
    def sum(self):
        self.result = np.sum(self.val)

spec1 = [('ListA',  nb.types.List(First.class_type.instance_type, reflected=True))]

@nb.jitclass(spec1)
class Combined:
    def __init__(self):
        self.ListA = [First(), First()] 
    def sum(self):
        for i, c in enumerate(self.ListA):
            c.sum()
    def getresult(self):
        result = []
        for i, c in enumerate(self.ListA):
            result.append(c.result)
        return result

C = Combined()
C.sum()
result = C.getresult()
print(result)

产生输出:[100.0, 100.0]

Intermezzo:在这里使用jitclass 有意义吗?

但是这里要记住的是,普通的 Python 类可能会比 jitclass-approach 快(或同样快):

import numpy as np
import numba as nb

class First:
    def __init__(self):
        self.type = 1
        self.val = np.ones(100)
        self.result = 0.
    def sum(self):
        self.result = np.sum(self.val)

class Combined:
    def __init__(self):
        self.ListA = [First(), First()] 
    def sum(self):
        for i, c in enumerate(self.ListA):
            c.sum()
    def getresult(self):
        result = []
        for i, c in enumerate(self.ListA):
            result.append(c.result)
        return result

C = Combined()
C.sum()
C.getresult()

如果只是出于好奇,那没问题。但是对于生产,我会从纯 Python+NumPy 开始,并且仅在速度太慢时才应用 numba,然后仅在成为瓶颈的部分且仅当 numba 擅长优化这些事情时才应用(numba 目前是专用工具,而不是通用工具)。

带有 numba 的异构(混合类型)列表?

在 no-python(无对象)模式下使用 numba,您需要同类列表。据我所知,numba 0.46 不支持在 jitclasses 或 nopython-jit 方法中包含不同类型对象的列表。这意味着您不能拥有一个包含 FirstSecond 实例的列表。

所以这是行不通的:

self.List.append(First())
self.List.append(Second())

来自numba docs

支持从 JIT 编译的函数以及所有方法和操作创建和返回列表。 列表必须是严格同质的:Numba 将拒绝任何包含不同类型对象的列表,即使这些类型是兼容的 [...]

【讨论】:

  • 不错。我很接近:)
  • 当您谈论同构列表时。这是否意味着我不能拥有不同大小的列表列表?就像一个 int 列表的列表:[ [1], [1, 2], [1, 2, 3] ]
  • @ymmx 列表列表对于 numba 来说有点棘手,但通常只要列表中的每个列表都包含相同的元素(大小无关紧要),它们就会受到支持。在您的情况下,这是正确的,因为每个子列表仅包含整数。但是,目前无法将此列表列表传递给 numba 函数(numba 错误,显示无法确定反射列表的反射列表的类型......)。但是我想知道为什么您会使用带有 numba 的列表列表。 numba 在列表方面并不比 Python 特别好。
猜你喜欢
  • 1970-01-01
  • 1970-01-01
  • 2016-12-05
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2020-01-13
  • 1970-01-01
相关资源
最近更新 更多