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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 16:15:06