在 Numba @jitclass 中使用列表
Using lists in Numba @jitclass
我正在模拟一款非常蹩脚的游戏,该游戏基本上会计算玩家在游戏中进行时收集的金币和敌人的数量。该代码包含两个 jitclasses:一个 player
jitclass 和一个 game
jitclass。
对于 player
class 我们有一些属性和一些方法来描述玩家在游戏中的进程。
from numba import jitclass, int64, float64, deferred_type
from numba.typed import List
import random
specs_player = OrderedDict()
specs_player['level'] = int64
specs_player['coins'] = float64
@jitclass(specs_player)
class Player:
def __init__(self):
self.level = 0
self.coins = 0
self.enemies = List()
def pass_level(self):
self.level += 1
def collect_coins(self, c):
self.coins += c
def collect_enemies(self, e):
self.enemies.append(e)
def reset_player(self):
self.level = 0
self.coins = 0
self.enemies = List()
如您所见,属性 enemies
是一个列表,它会随着玩家在游戏中的进行而附加值。
game
jitclass 使用前两行调用 player
jitclass 作为属性:
Player_type = deferred_type()
Player_type.define(Player.class_type.instance_type)
specs_Game = OrderedDict()
specs_Game['last_level'] = int64
specs_Game['diff_threshold'] = float64
specs_Game['player'] = Player_type
class Game:
def __init__(self, l, t):
self.player = Player()
self.last_level = l
self.diff_threshold = t
def play_gameround(self):
random_draw = random.uniform(0, 1)
if random_draw > self.diff_threshold:
# Pass Level
self.player.pass_level()
# Collect coins
coins_earned = 100*(random_draw - self.diff_threshold)
self.player.collect_coins(coins_earned)
#Collect enemies
if coins_earned > 10:
self.player.collect_enemies()
def reset_gameplay(self):
self.player.reset_player()
def continue_playing(self):
condition = self.player.level < self.last_level
return(condition)
最后,一个名为 run_one_player
的函数模拟一个玩家的进程和 returns 一个包含所有必要数据的数组:
def run_one_player(gameplay):
while gameplay.continue_playing():
gameplay.play_gameround()
player_data = np.array([gameplay.player.level,
gameplay.player.coins,
gameplay.player.enemies])
return (player_data)
到运行我刚刚输入的代码:
g = Game(l = 10, t = 0.5)
data_list = run_one_player(g)
但是,这不起作用,Numba returns 出现以下错误,我很确定这是因为我没有为 enemies
字段正确定义 Numba 类型.
---------------------------------------------------------------------------
TypingError Traceback (most recent call last)
<ipython-input-51-f6340d9cfcd7> in <module>
1 players = 10
----> 2 g = Game(l = 10, t = 0.5)
3
4 data_list = run_one_player(g)
<ipython-input-49-6e42f0104cf1> in __init__(self, l, t)
46
47 def __init__(self, l, t):
---> 48 self.player = Player()
49 self.last_level = l
50 self.diff_threshold = t
/Library/Frameworks/Python.framework/Versions/3.7/lib/python3.7/site-packages/numba/jitclass/base.py in __call__(cls, *args, **kwargs)
124 bind = cls._ctor_sig.bind(None, *args, **kwargs)
125 bind.apply_defaults()
--> 126 return cls._ctor(*bind.args[1:], **bind.kwargs)
127
128
/Library/Frameworks/Python.framework/Versions/3.7/lib/python3.7/site-packages/numba/dispatcher.py in _compile_for_args(self, *args, **kws)
374 e.patch_message(msg)
375
--> 376 error_rewrite(e, 'typing')
377 except errors.UnsupportedError as e:
378 # Something unsupported is present in the user code, add help info
/Library/Frameworks/Python.framework/Versions/3.7/lib/python3.7/site-packages/numba/dispatcher.py in error_rewrite(e, issue_type)
341 raise e
342 else:
--> 343 reraise(type(e), e, None)
344
345 argtypes = []
/Library/Frameworks/Python.framework/Versions/3.7/lib/python3.7/site-packages/numba/six.py in reraise(tp, value, tb)
656 value = tp()
657 if value.__traceback__ is not tb:
--> 658 raise value.with_traceback(tb)
659 raise value
660
TypingError: Failed in nopython mode pipeline (step: nopython frontend)
Failed in nopython mode pipeline (step: nopython frontend)
Cannot resolve setattr: (instance.jitclass.Player#123b442d0<level:int64,coins:float64>).enemies = ListType[undefined]
File "<ipython-input-49-6e42f0104cf1>", line 18:
def __init__(self):
<source elided>
self.coins = 0
self.enemies = List()
^
[1] During: typing of set attribute 'enemies' at <ipython-input-49-6e42f0104cf1> (18)
File "<ipython-input-49-6e42f0104cf1>", line 18:
def __init__(self):
<source elided>
self.coins = 0
self.enemies = List()
^
[1] During: resolving callee type: jitclass.Player#123b442d0<level:int64,coins:float64>
[2] During: typing of call at <string> (3)
[3] During: resolving callee type: jitclass.Player#123b442d0<level:int64,coins:float64>
[4] During: typing of call at <string> (3)
File "<string>", line 3:
<source missing, REPL/exec in use?>
This is not usually a problem with Numba itself but instead often caused by
the use of unsupported features or an issue in resolving types.
To see Python/NumPy features supported by the latest release of Numba visit:
http://numba.pydata.org/numba-doc/latest/reference/pysupported.html
and
http://numba.pydata.org/numba-doc/latest/reference/numpysupported.html
For more information about typing errors and how to debug them visit:
http://numba.pydata.org/numba-doc/latest/user/troubleshoot.html#my-code-doesn-t-compile
If you think your code should work with Numba, please report the error message
and traceback, along with a minimal reproducer at:
https://github.com/numba/numba/issues/new
首先:我认为您不应该将 numba 用于这样的事情。 Numba 是一种专门的工具,非常擅长解决特定类型的问题,而这不是其中之一:
1.1.2. Will Numba work for my code?
This depends on what your code looks like, if your code is numerically orientated (does a lot of math), uses NumPy a lot and/or has a lot of loops, then Numba is often a good choice
[...]
然而,在您的特定情况下,您需要完整键入 jitclass
的 所有 属性。这意味着您必须使用 numba 理解的类型(支持的类型之一或另一个 jitclass)输入 enemies
,否则它将无法工作。
由于您没有提供类型,我们假设它是一个整数:
import numba as nb
specs_player = {}
specs_player['level'] = nb.int64
specs_player['coins'] = nb.float64
specs_player['enemies'] = nb.types.List(nb.int64)
@nb.jitclass(specs_player)
class Player:
def __init__(self):
self.level = 0
self.coins = 0
self.enemies = []
创建新实例时失败,因为 numba 无法推断空列表的类型(至少目前是这样)。所以你必须用某种类型进行初始化。我没有找到比创建包含项目的列表然后清除它更好的方法:
import numba as nb
specs_player = {}
specs_player['level'] = nb.int64
specs_player['coins'] = nb.float64
specs_player['enemies'] = nb.types.List(nb.int64)
@nb.njit
def empty_int64_list():
l = [nb.int64(10)]
l.clear()
return l
@nb.jitclass(specs_player)
class Player:
def __init__(self):
self.level = 0
self.coins = 0
self.enemies = empty_int64_list()
如果您的 enemies
不是整数,情况可能会复杂得多。但是,我认为在您的情况下这不值得,因为这不是 numba 比纯 Python.
更有效(显着)解决的问题
我正在模拟一款非常蹩脚的游戏,该游戏基本上会计算玩家在游戏中进行时收集的金币和敌人的数量。该代码包含两个 jitclasses:一个 player
jitclass 和一个 game
jitclass。
对于 player
class 我们有一些属性和一些方法来描述玩家在游戏中的进程。
from numba import jitclass, int64, float64, deferred_type
from numba.typed import List
import random
specs_player = OrderedDict()
specs_player['level'] = int64
specs_player['coins'] = float64
@jitclass(specs_player)
class Player:
def __init__(self):
self.level = 0
self.coins = 0
self.enemies = List()
def pass_level(self):
self.level += 1
def collect_coins(self, c):
self.coins += c
def collect_enemies(self, e):
self.enemies.append(e)
def reset_player(self):
self.level = 0
self.coins = 0
self.enemies = List()
如您所见,属性 enemies
是一个列表,它会随着玩家在游戏中的进行而附加值。
game
jitclass 使用前两行调用 player
jitclass 作为属性:
Player_type = deferred_type()
Player_type.define(Player.class_type.instance_type)
specs_Game = OrderedDict()
specs_Game['last_level'] = int64
specs_Game['diff_threshold'] = float64
specs_Game['player'] = Player_type
class Game:
def __init__(self, l, t):
self.player = Player()
self.last_level = l
self.diff_threshold = t
def play_gameround(self):
random_draw = random.uniform(0, 1)
if random_draw > self.diff_threshold:
# Pass Level
self.player.pass_level()
# Collect coins
coins_earned = 100*(random_draw - self.diff_threshold)
self.player.collect_coins(coins_earned)
#Collect enemies
if coins_earned > 10:
self.player.collect_enemies()
def reset_gameplay(self):
self.player.reset_player()
def continue_playing(self):
condition = self.player.level < self.last_level
return(condition)
最后,一个名为 run_one_player
的函数模拟一个玩家的进程和 returns 一个包含所有必要数据的数组:
def run_one_player(gameplay):
while gameplay.continue_playing():
gameplay.play_gameround()
player_data = np.array([gameplay.player.level,
gameplay.player.coins,
gameplay.player.enemies])
return (player_data)
到运行我刚刚输入的代码:
g = Game(l = 10, t = 0.5)
data_list = run_one_player(g)
但是,这不起作用,Numba returns 出现以下错误,我很确定这是因为我没有为 enemies
字段正确定义 Numba 类型.
---------------------------------------------------------------------------
TypingError Traceback (most recent call last)
<ipython-input-51-f6340d9cfcd7> in <module>
1 players = 10
----> 2 g = Game(l = 10, t = 0.5)
3
4 data_list = run_one_player(g)
<ipython-input-49-6e42f0104cf1> in __init__(self, l, t)
46
47 def __init__(self, l, t):
---> 48 self.player = Player()
49 self.last_level = l
50 self.diff_threshold = t
/Library/Frameworks/Python.framework/Versions/3.7/lib/python3.7/site-packages/numba/jitclass/base.py in __call__(cls, *args, **kwargs)
124 bind = cls._ctor_sig.bind(None, *args, **kwargs)
125 bind.apply_defaults()
--> 126 return cls._ctor(*bind.args[1:], **bind.kwargs)
127
128
/Library/Frameworks/Python.framework/Versions/3.7/lib/python3.7/site-packages/numba/dispatcher.py in _compile_for_args(self, *args, **kws)
374 e.patch_message(msg)
375
--> 376 error_rewrite(e, 'typing')
377 except errors.UnsupportedError as e:
378 # Something unsupported is present in the user code, add help info
/Library/Frameworks/Python.framework/Versions/3.7/lib/python3.7/site-packages/numba/dispatcher.py in error_rewrite(e, issue_type)
341 raise e
342 else:
--> 343 reraise(type(e), e, None)
344
345 argtypes = []
/Library/Frameworks/Python.framework/Versions/3.7/lib/python3.7/site-packages/numba/six.py in reraise(tp, value, tb)
656 value = tp()
657 if value.__traceback__ is not tb:
--> 658 raise value.with_traceback(tb)
659 raise value
660
TypingError: Failed in nopython mode pipeline (step: nopython frontend)
Failed in nopython mode pipeline (step: nopython frontend)
Cannot resolve setattr: (instance.jitclass.Player#123b442d0<level:int64,coins:float64>).enemies = ListType[undefined]
File "<ipython-input-49-6e42f0104cf1>", line 18:
def __init__(self):
<source elided>
self.coins = 0
self.enemies = List()
^
[1] During: typing of set attribute 'enemies' at <ipython-input-49-6e42f0104cf1> (18)
File "<ipython-input-49-6e42f0104cf1>", line 18:
def __init__(self):
<source elided>
self.coins = 0
self.enemies = List()
^
[1] During: resolving callee type: jitclass.Player#123b442d0<level:int64,coins:float64>
[2] During: typing of call at <string> (3)
[3] During: resolving callee type: jitclass.Player#123b442d0<level:int64,coins:float64>
[4] During: typing of call at <string> (3)
File "<string>", line 3:
<source missing, REPL/exec in use?>
This is not usually a problem with Numba itself but instead often caused by
the use of unsupported features or an issue in resolving types.
To see Python/NumPy features supported by the latest release of Numba visit:
http://numba.pydata.org/numba-doc/latest/reference/pysupported.html
and
http://numba.pydata.org/numba-doc/latest/reference/numpysupported.html
For more information about typing errors and how to debug them visit:
http://numba.pydata.org/numba-doc/latest/user/troubleshoot.html#my-code-doesn-t-compile
If you think your code should work with Numba, please report the error message
and traceback, along with a minimal reproducer at:
https://github.com/numba/numba/issues/new
首先:我认为您不应该将 numba 用于这样的事情。 Numba 是一种专门的工具,非常擅长解决特定类型的问题,而这不是其中之一:
1.1.2. Will Numba work for my code?
This depends on what your code looks like, if your code is numerically orientated (does a lot of math), uses NumPy a lot and/or has a lot of loops, then Numba is often a good choice
[...]
然而,在您的特定情况下,您需要完整键入 jitclass
的 所有 属性。这意味着您必须使用 numba 理解的类型(支持的类型之一或另一个 jitclass)输入 enemies
,否则它将无法工作。
由于您没有提供类型,我们假设它是一个整数:
import numba as nb
specs_player = {}
specs_player['level'] = nb.int64
specs_player['coins'] = nb.float64
specs_player['enemies'] = nb.types.List(nb.int64)
@nb.jitclass(specs_player)
class Player:
def __init__(self):
self.level = 0
self.coins = 0
self.enemies = []
创建新实例时失败,因为 numba 无法推断空列表的类型(至少目前是这样)。所以你必须用某种类型进行初始化。我没有找到比创建包含项目的列表然后清除它更好的方法:
import numba as nb
specs_player = {}
specs_player['level'] = nb.int64
specs_player['coins'] = nb.float64
specs_player['enemies'] = nb.types.List(nb.int64)
@nb.njit
def empty_int64_list():
l = [nb.int64(10)]
l.clear()
return l
@nb.jitclass(specs_player)
class Player:
def __init__(self):
self.level = 0
self.coins = 0
self.enemies = empty_int64_list()
如果您的 enemies
不是整数,情况可能会复杂得多。但是,我认为在您的情况下这不值得,因为这不是 numba 比纯 Python.