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

在Numba即时编译函数中拼接NumPy数组的问题

在Numba njit函数中简洁拼接NumPy数组的方案

在Numba即时编译(njit)函数中,直接用np.hstack处理Python原生数组列表/元组会编译失败——这是因为Numba无法提前推断动态容器内元素的类型和形状信息。你自己实现的join_arrays_c虽然可行,但代码较为繁琐,以下是两种更简洁的替代方案:

方案1:使用Numba类型化列表(numba.typed.List)+ np.hstack

Numba的类型化容器能明确标注元素类型,让np.hstack可以正常编译执行,写法最接近原生NumPy的简洁风格:

import numpy as np
from numba import njit
from numba.typed import List

@njit
def join_arrays_typed(list_of_arrays):
    return np.hstack(list_of_arrays)

@njit
def my_program():
    array_1 = np.array([0,3])
    array_2 = np.array([0,4,2,3])
    array_3 = np.array([9,1,3,3,5,9])
    
    # 创建Numba类型化列表并添加数组
    typed_list = List()
    typed_list.append(array_1)
    typed_list.append(array_2)
    typed_list.append(array_3)
    
    return join_arrays_typed(typed_list)

print(my_program())  # 输出: [0 3 0 4 2 3 9 1 3 3 5 9]

方案2:简化手动拼接代码

如果不想引入类型化列表,可以优化手动拼接的逻辑,用切片赋值替代内层循环,代码更紧凑:

import numpy as np
from numba import njit

@njit
def join_arrays_simplified(arrays):
    tot_len = sum(len(arr) for arr in arrays)
    new_array = np.zeros(tot_len, dtype=np.int64)
    idx = 0
    for arr in arrays:
        arr_len = len(arr)
        # 用切片赋值替换内层循环
        new_array[idx:idx+arr_len] = arr
        idx += arr_len
    return new_array

@njit
def my_program():
    array_1 = np.array([0,3])
    array_2 = np.array([0,4,2,3])
    array_3 = np.array([9,1,3,3,5,9])
    
    list_of_arrays = [array_1, array_2, array_3]
    
    return join_arrays_simplified(list_of_arrays)

print(my_program())  # 输出: [0 3 0 4 2 3 9 1 3 3 5 9]

说明

  • 方案1的优势是写法简洁,和原生NumPy的使用习惯一致,但需要额外创建Numba类型化列表;
  • 方案2的代码比你原来的join_arrays_c更紧凑,同时保持了对原生Python列表的兼容性,无需额外依赖。

内容的提问来源于stack exchange,提问作者Patrick

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 02:17:42