Numba传入字典地址参数时报错,如何正确传递该类参数?
问题原因与解决办法
问题根源
Numba的@jit(nopython=True)模式不支持Python的可变关键字参数(即**kwargs语法)。你定义的函数用**shape接收关键字参数,但通过random(**shape)解包字典时,相当于传递了0=100和1=100两个关键字参数,而Numba编译后的函数无法识别这种参数形式,因此抛出参数数量不匹配的错误。
解决办法
方案1:改用位置参数,直接传递形状元组
修改函数定义为接收明确的位置参数,将形状以元组形式传递(或解包传递):
from timeit import default_timer as timer from datetime import timedelta import numpy as np from numba import jit @jit(nopython=True, parallel=True) def random(rows, cols): return np.random.random((rows, cols)) shape = (100, 100) start = timer() random(*shape) # 解包元组传递参数 end = timer() print(timedelta(seconds=end-start))
方案2:保留字典,提取值转为元组后传递
如果必须用字典存储形状,可提取字典的值转为元组,再解包传递给函数:
from timeit import default_timer as timer from datetime import timedelta import numpy as np from numba import jit @jit(nopython=True, parallel=True) def random(rows, cols): return np.random.random((rows, cols)) shape = {0:100, 1:100} start = timer() random(*shape.values()) # 提取字典值并解包 end = timer() print(timedelta(seconds=end-start))
方案3:直接接收形状元组作为单个参数
让函数接收一个形状元组参数,传递更简洁:
from timeit import default_timer as timer from datetime import timedelta import numpy as np from numba import jit @jit(nopython=True, parallel=True) def random(shape): return np.random.random(shape) shape = (100, 100) # 从字典转换的话:shape = tuple(shape_dict.values()) start = timer() random(shape) end = timer() print(timedelta(seconds=end-start))
内容的提问来源于stack exchange,提问作者DuttaA
相关产品推荐
相关产品推荐

