Numpy如何对现有数组进行可重复索引的向量化修改?
解决Numpy数组重复索引的累加修改问题
当使用花式索引对Numpy数组的同一索引进行多次累加时,直接用+=操作会因为返回副本而非视图,导致重复索引的修改只生效最后一次。Numpy提供了原生方法解决这个问题:
方法1:使用np.add.at原地操作
np.add.at是专门针对重复索引设计的原地累加函数,会对每个索引位置执行多次累加:
import numpy as np zeros = np.zeros(10) indices = np.array([0,0]) adders = np.array([5,8]) np.add.at(zeros, indices, adders) print(zeros) # 输出:array([13., 0., 0., 0., 0., 0., 0., 0., 0., 0.])
方法2:使用np.bincount统计总增量后赋值
先通过np.bincount计算每个索引对应的总增量,再一次性加到原数组上:
import numpy as np zeros = np.zeros(10) indices = np.array([0,0]) adders = np.array([5,8]) total_increments = np.bincount(indices, weights=adders, minlength=len(zeros)) zeros += total_increments print(zeros) # 输出:array([13., 0., 0., 0., 0., 0., 0., 0., 0., 0.])
这两种方法都无需手动循环,完全利用Numpy原生功能实现重复索引的多次修改。
内容的提问来源于stack exchange,提问作者Estif
相关产品推荐
相关产品推荐

