如何在Numba中创建存储异构元组值的空TypedDict?
问题:初始化Numba类型字典(字符串键+异构元组值)失败
我想要初始化一个以字符串为键、异构元组为值的Numba类型字典,实现代码如下:
from numba import types, typed, njit, typeof, from_dtype TypeMyTuple = types.Tuple([types.unicode_type, types.float16, types.int16]) my_dict = typed.Dict.empty(key_type=types.unicode_type, value_type=TypeMyTuple)
运行后出现如下错误:
Traceback (most recent call last): File "/home/me/.pycharm_helpers/pydev/pydevconsole.py", line 364, in runcode coro = func() ^^^^^^ File "<input>", line 8, in <module> File "/home/me/.pyenv/versions/3.11.1/envs/my_env/lib/python3.11/site-packages/numba/typed/typeddict.py", line 105, in empty return cls(dcttype=DictType(key_type, value_type), n_keys=n_keys) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/home/me/.pyenv/versions/3.11.1/envs/my_env/lib/python3.11/site-packages/numba/typed/typeddict.py", line 120, in __init__ self._dict_type, self._opaque = self._parse_arg(**kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^ File "/home/me/.pyenv/versions/3.11.1/envs/my_env/lib/python3.11/site-packages/numba/typed/typeddict.py", line 156, in _parse_arg opaque = _make_dict(dcttype.key_type, dcttype.value_type, ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/home/me/.pyenv/versions/3.11.1/envs/my_env/lib/python3.11/site-packages/numba/core/dispatcher.py", line 468, in _compile_for_args error_rewrite(e, 'typing') File "/home/me/.pyenv/versions/3.11.1/envs/my_env/lib/python3.11/site-packages/numba/core/dispatcher.py", line 409, in error_rewrite raise e.with_traceback(None) numba.core.errors.TypingError: Failed in nopython mode pipeline (step: nopython frontend) No implementation of function Function(<function new_dict at 0x7f78a6293920>) found for signature: >>> new_dict(typeref[unicode_type], typeref[Tuple(unicode_type, float16, int16)], n_keys=int64) There are 2 candidate implementations: - Of which 2 did not match due to: Overload in function 'impl_new_dict': File: numba/typed/dictobject.py: Line 653. With argument(s): '(typeref[unicode_type], typeref[Tuple(unicode_type, float16, int16)], n_keys=int64)': Rejected as the implementation raised a specific error: LoweringError: Failed in nopython mode pipeline (step: native lowering) float16 File "../../../../../../home/me/.pyenv/versions/3.11.1/envs/my_env/lib/python3.11/site-packages/numba/typed/dictobject.py", line 671: def imp(key, value, n_keys=0): <source elided> raise RuntimeError("expecting *n_keys* to be >= 0") dp = _dict_new_sized(n_keys, keyty, valty) ^ During: lowering "dp = call $48load_global.0(n_keys, $62load_deref.3, $64load_deref.4, func=$48load_global.0, args=[Var(n_keys, dictobject.py:668), Var($62load_deref.3, dictobject.py:671), Var($64load_deref.4, dictobject.py:671)], kws=(), vararg=None, varkwarg=None, target=None)" at /home/me/.pyenv/versions/3.11.1/envs/my_env/lib/python3.11/site-packages/numba/typed/dictobject.py (671) raised from /home/me/.pyenv/versions/3.11.1/envs/my_env/lib/python3.11/site-packages/numba/core/errors.py:846 During: resolving callee type: Function(<function new_dict at 0x7f78a6293920>) During: typing of call at /home/me/.pyenv/versions/3.11.1/envs/my_env/lib/python3.11/site-packages/numba/typed/typeddict.py (23) File "../../../../../../home/me/.pyenv/versions/3.11.1/envs/my_env/lib/python3.11/site-packages/numba/typed/typeddict.py", line 23: def _make_dict(keyty, valty, n_keys=0): return dictobject._as_meminfo(dictobject.new_dict(keyty, valty, ^
解决思路
从错误信息核心来看,问题出在float16类型在Numba typed字典的元组值中存在底层实现支持缺失,可尝试以下几种方案:
替换
float16为更高精度浮点类型:
将元组中的types.float16替换为types.float32或types.float64,这两种类型在Numba的typed容器中支持更完善。修改后代码示例:from numba import types, typed TypeMyTuple = types.Tuple([types.unicode_type, types.float32, types.int16]) my_dict = typed.Dict.empty(key_type=types.unicode_type, value_type=TypeMyTuple)在
njit函数内部初始化字典:
不在全局直接创建typed字典,而是在njit装饰的函数内部初始化并操作。Numba在nopython模式下对元组类型的处理更灵活,示例:from numba import types, typed, njit @njit def create_dict(): my_dict = typed.Dict.empty(types.unicode_type, types.Tuple((types.unicode_type, types.float16, types.int16))) # 可在此添加键值对 my_dict["test"] = ("abc", 1.5, 10) return my_dict my_dict = create_dict()用自定义结构体替代异构元组:
如果必须使用float16,可以定义Numba自定义结构体类型替代元组,结构体在底层支持更稳定:from numba import types, typed, njit # 定义结构体类型 MyStruct = types.Record.make_struct([ ("str_field", types.unicode_type), ("float_field", types.float16), ("int_field", types.int16) ]) @njit def create_struct_dict(): my_dict = typed.Dict.empty(types.unicode_type, MyStruct) # 构造结构体实例并添加 struct_val = MyStruct("abc", 1.5, 10) my_dict["test"] = struct_val return my_dict my_dict = create_struct_dict()
内容的提问来源于stack exchange,提问作者GlaceCelery
相关产品推荐
相关产品推荐

