如何在Numpy中高效实现类似np.put的数组指定位置累加求和(无循环)
高效实现类似numpy.put的累加操作(而非替换)
嘿,这个需求完全可以用NumPy的原生矢量化操作高效解决,根本不需要写循环!下面给你两种实用的方案,都是性能拉满的选择:
方法一:用np.add.at(最直接推荐)
np.add.at是NumPy专门用来处理原地累加的工具,完美适配你的需求——它会在指定索引位置上累加对应的值,而不是像np.put那样直接替换。而且不管索引有没有重复,都能正确处理。
举个和你例子匹配的代码:
import numpy as np # 初始化你的数组 a = np.array([[4, 4], [2, 0]]) indices = np.array([0, 1]) # 对应扁平化后的数组索引 b = np.array([5, 6]) # 执行累加操作 np.add.at(a.ravel(), indices, b) print(a) # 输出结果:[[ 9 10] # [ 2 0]]
这里a.ravel()把二维数组扁平化,np.add.at会找到indices对应的位置,把b里的元素逐个加到a的对应位置上,完全符合你的期望。
方法二:用np.bincount(适合需要额外统计的场景)
如果你需要对索引的权重做额外处理,np.bincount也是个不错的选择。它会先统计每个索引对应的累加值,再加到原数组上:
import numpy as np a = np.array([[4, 4], [2, 0]]) indices = np.array([0, 1]) b = np.array([5, 6]) flat_a = a.ravel() # 计算每个索引的累加权重 counts = np.bincount(indices, weights=b, minlength=flat_a.size) # 加到原数组 flat_a += counts # 恢复原形状 a = flat_a.reshape(a.shape) print(a) # 同样得到期望结果
这种方法在处理大量重复索引时也很高效,不过代码比第一种稍长一些。
为什么不用循环?
NumPy的矢量化操作是底层用C实现的,比Python循环快几个数量级,尤其是当数组规模很大的时候,性能差距会非常明显。所以优先用这些原生方法就对了!
内容的提问来源于stack exchange,提问作者rvinas
相关产品推荐
相关产品推荐

