如何在Numba中实现np.repeat的axis=2功能(沿最后维度重复数组)
Numba兼容的沿最后维度重复数组实现方案
你的需求是将形状为(1000,1)的数组扩展为(1000,1,1)后,沿最后一维重复10次得到(1000,1,10)的数组,但Numba不支持np.repeat的axis参数,也不允许用列表作为np.dstack的输入。以下是两种高效的兼容实现:
方案一:利用广播赋值(针对最后维度为1的场景)
这是最高效的方式,利用Numba支持的切片广播赋值特性,无需显式循环:
import numba import numpy as np original = np.random.rand(1000,1) no_repeats = 10 @numba.jit(nopython=True) def repeat_last_dim_broadcast(arr): # arr 输入形状为 (N, 1, 1) out_shape = arr.shape[:-1] + (no_repeats,) out = np.empty(out_shape, dtype=arr.dtype) # 利用广播将arr的值填充到out的所有最后一维位置 out[:] = arr return out # 先预处理得到扩展后的数组 expanded_original = 5 * original[:, :, np.newaxis] big_original = repeat_last_dim_broadcast(expanded_original)
方案二:通用手动循环(支持任意最后维度的重复)
如果你的场景需要对任意最后维度长度的数组进行重复,可以用Numba优化的嵌套循环,性能接近原生numpy:
@numba.jit(nopython=True) def repeat_last_dim_general(arr, no_repeats): # 获取输入数组的形状 *pre_dims, last_dim = arr.shape out_shape = tuple(pre_dims) + (last_dim * no_repeats,) out = np.empty(out_shape, dtype=arr.dtype) # 遍历前n-1维,对每个元素在最后维度重复 for idx in np.ndindex(*pre_dims): for i in range(last_dim): val = arr[idx + (i,)] for r in range(no_repeats): out[idx + (i*no_repeats + r,)] = val return out # 使用示例 expanded_original = 5 * original[:, :, np.newaxis] big_original = repeat_last_dim_general(expanded_original, no_repeats)
补充说明
- 方案一仅适用于最后维度长度为1的情况,因为广播赋值会自动将单元素维度扩展到目标长度;
- 方案二更通用,支持任意最后维度的重复操作,Numba会将循环编译为高效的机器码,性能不会比原生numpy差;
- 避免使用Python列表作为numpy函数的输入,Numba的nopython模式完全不支持Python容器类型,所有操作都要基于numpy数组或标量。
内容的提问来源于stack exchange,提问作者codephantom
相关产品推荐
相关产品推荐

