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

如何快速从同形状Numpy数组筛选元素并计算横向和?

优化Numpy数组按条件分组行求和的效率

问题背景

有两个同形状的二维Numpy数组A(代码中为varr)和B(代码中为rarr),需要针对digits中的每个值x,提取B中等于x的位置对应的A元素,计算每行的求和结果。现有两种实现因数组规模较大速度偏慢,需要更高效的方案。

原方法一:

farr = np.stack([np.where(rarr == d, varr, 0).sum(axis=1) for d in digits])

原方法二(Numba加速):

@njit
def final_array(varr, rarr, digits):
  farr = np.zeros((len(varr), len(digits)))
  for i, d in enumerate(digits):
    farr[:,i] = np.where(rarr == d, varr, 0).sum(axis=1)
  return farr

farr = final_array(varr, rarr, digits)

核心问题分析

原两种方法本质都是循环digits中的每个值,每次遍历整个数组做条件判断和求和,时间复杂度为O(K*N*M)(K是digits长度,N、M是数组的行、列数),当K或数组规模较大时,重复遍历的开销会非常明显。

优化方案

方案1:广播+向量化计算(无循环,代码简洁)

利用Numpy的广播特性,一次性完成所有digits的条件判断与求和,仅遍历数组一次:

# 将digits转为列向量,实现与二维数组的广播匹配
digits_col = digits.reshape(-1, 1, 1)
# 生成形状为(K, N, M)的布尔掩码,标记B中等于对应digit的位置
mask = (rarr == digits_col)
# 掩码与A相乘后按列求和(对应原需求的横向求和),得到(K, N)的结果
farr = (mask * varr).sum(axis=2)

# 若需要和方法二一致的(N, K)形状,只需转置
# farr = farr.T

方案2:np.bincount分组求和(适合digits为连续整数场景)

如果digits是连续的整数范围,np.bincount能以接近O(N*M)的时间复杂度完成计算,效率极高:

N, M = varr.shape
# 生成每行的索引,将二维位置映射为一维的(row_idx * max_digit + digit)格式
row_indices = np.repeat(np.arange(N), M)
max_digit = max(digits)
# 按(row_idx, digit)分组求和
counts = np.bincount(
    row_indices * (max_digit + 1) + rarr.ravel(),
    weights=varr.ravel(),
    minlength=N * (max_digit + 1)
)
# 重塑形状后提取digits对应的列
farr = counts.reshape(N, -1)[:, digits]

方案3:Numba单遍历优化(避免重复数组扫描)

原Numba方法仍重复遍历数组K次,改为仅遍历数组一次,同时更新对应digit的行和:

from numba import njit

@njit
def final_array_fast(varr, rarr, digits):
    N, M = varr.shape
    K = len(digits)
    # 预构建digit到结果列索引的映射,避免循环内查找
    digit_map = {d: idx for idx, d in enumerate(digits)}
    farr = np.zeros((N, K))
    
    # 仅遍历数组一次,累加符合条件的元素
    for i in range(N):
        for j in range(M):
            current_digit = rarr[i, j]
            if current_digit in digit_map:
                farr[i, digit_map[current_digit]] += varr[i, j]
    return farr

farr = final_array_fast(varr, rarr, digits)

方案选择建议

  • 若digits数量不多、数组维度适中,优先选方案1,代码简洁且完全向量化,无需额外依赖。
  • 若digits是连续整数,方案2的效率最高,时间复杂度接近线性。
  • 若数组规模极大且digits数量多,方案3的单遍历Numba实现能最小化内存访问开销,速度最优。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 16:45:33