Numba njit多输入输出函数签名编写:处理双类型返回值
关于Numba njit函数支持多返回类型及签名编写的问题
原代码
import numpy as np import numba as nb sorted_fb_i = np.array([1, 3, 4, 2, 5], np.int64) fb_groups_ids = nb.typed.List([np.array([4, 2], np.int64), np.array([1, 3, 5], np.int64)]) moved_fb_group_ids = nb.typed.List.empty_list(nb.types.Array(dtype=nb.int64, ndim=1, layout="C")) ind = 0 @nb.njit def points_group_ids(sorted_fb_i, fb_groups_ids, moved_fb_group_ids, ind): pnt_group_ids_ = sorted_fb_i[ind] for i in range(len(fb_groups_ids)): if sorted_fb_i[ind] in fb_groups_ids[i]: pnt_group_ids_ = fb_groups_ids[i] moved_fb_group_ids.append(fb_groups_ids.pop(i)) break return pnt_group_ids_, fb_groups_ids, moved_fb_group_ids
报错信息
Cannot unify array(int64, 1d, C) and int64 for 'pnt_group_ids_.2'
问题与解答
1. 能否编写支持两种返回类型的函数签名?
Numba的njit基于静态类型编译,要求函数内变量类型在编译阶段固定,不允许动态切换(一会是int64,一会是int64[::1])。直接通过签名让同一个变量兼容两种类型不可行——这也是你遇到类型无法统一报错的核心原因。
如果尝试用联合类型签名,写法如下,但大概率仍无法通过编译(因为Numba数据流分析不允许单变量在函数内改变类型):
import numpy as np import numba as nb # 定义类型 int64_type = nb.int64 int64_array_type = nb.types.Array(nb.int64, 1, 'C') union_return_type = nb.types.Union((int64_type, int64_array_type)) list_of_arrays_type = nb.types.ListType(int64_array_type) @nb.njit( signature_or_function=(int64_array_type, list_of_arrays_type, list_of_arrays_type, int64_type), return_type=(union_return_type, list_of_arrays_type, list_of_arrays_type) ) def points_group_ids(sorted_fb_i, fb_groups_ids, moved_fb_group_ids, ind): pnt_group_ids_ = sorted_fb_i[ind] for i in range(len(fb_groups_ids)): if sorted_fb_i[ind] in fb_groups_ids[i]: pnt_group_ids_ = fb_groups_ids[i] moved_fb_group_ids.append(fb_groups_ids.pop(i)) break return pnt_group_ids_, fb_groups_ids, moved_fb_group_ids
2. 修改代码统一返回类型后,如何编写多输入多输出的签名?
将pnt_group_ids_ = sorted_fb_i[ind]改为pnt_group_ids_ = np.array([sorted_fb_i[ind]], np.int64)后,返回类型统一为数组,此时可以正确编写签名。你之前遇到的TypeError: 'tuple' object is not callable是因为签名写法错误,正确写法如下:
方式一:通过类型对象定义
import numpy as np import numba as nb # 定义类型 int64_array = nb.types.Array(nb.int64, 1, 'C') list_of_int64_arrays = nb.types.ListType(int64_array) @nb.njit((int64_array, list_of_int64_arrays, list_of_int64_arrays, nb.int64)) def points_group_ids(sorted_fb_i, fb_groups_ids, moved_fb_group_ids, ind): pnt_group_ids_ = np.array([sorted_fb_i[ind]], np.int64) for i in range(len(fb_groups_ids)): if sorted_fb_i[ind] in fb_groups_ids[i]: pnt_group_ids_ = fb_groups_ids[i] moved_fb_group_ids.append(fb_groups_ids.pop(i)) break return pnt_group_ids_, fb_groups_ids, moved_fb_group_ids
方式二:通过类型字符串定义
@nb.njit("(int64[::1], ListType(int64[::1]), ListType(int64[::1]), int64)") def points_group_ids(sorted_fb_i, fb_groups_ids, moved_fb_group_ids, ind): pnt_group_ids_ = np.array([sorted_fb_i[ind]], np.int64) for i in range(len(fb_groups_ids)): if sorted_fb_i[ind] in fb_groups_ids[i]: pnt_group_ids_ = fb_groups_ids[i] moved_fb_group_ids.append(fb_groups_ids.pop(i)) break return pnt_group_ids_, fb_groups_ids, moved_fb_group_ids
3. fb_groups_ids为空时是否会报错?
不会报错。当fb_groups_ids为空时,len(fb_groups_ids)返回0,range(0)不会产生迭代,循环体不执行,pnt_group_ids_保持初始的单元素数组,函数可正常返回。
核心需求的最优方案
优先不修改代码支持两种返回类型的方案不可行,因为Numba静态类型机制不允许变量在函数内动态切换类型。最优方案是修改代码统一返回类型(改为单元素数组),再按上述正确方式编写函数签名,既能通过编译,也能满足业务逻辑需求。
内容的提问来源于stack exchange,提问作者Ali_Sh
相关产品推荐
相关产品推荐

