You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.16 14:30:58