如何对NumPy数组按重复索引累加对应值得到目标结果?
最简实现方式:使用
np.add.at()完成重复索引累加 你遇到的问题其实是NumPy花式索引里的典型特性:当用重复索引做原地赋值时,底层是先创建视图再批量赋值,重复的索引位置会被最后一次的计算值覆盖,而非逐个累加。要实现你想要的「遍历索引列表,把对应值依次累加到目标位置」的逻辑,np.add.at()就是专门解决这个场景的最简方案。
直接看完整代码示例:
import numpy as np # 初始化数组 a = np.array([0, 0]) indices = [0, 0, 1, 1] values_to_add = [1, 2, 3, 4] # 执行重复索引累加 np.add.at(a, indices, values_to_add) print(a) # 输出:array([3, 7])
为什么这个方法可行?
np.add.at()是NumPy提供的原地操作函数,它会严格按照索引列表的顺序,将对应的值逐个累加到目标数组的指定索引位置,完全不会出现重复索引覆盖的问题。对应你的需求:
- 索引0会依次加上1和2,最终计算为0+1+2=3
- 索引1会依次加上3和4,最终计算为0+3+4=7
而你原来用a[[0,0,1,1]] += [1,2,3,4]得到[2,4],本质是因为这个操作的执行逻辑:
- 先从
a中取出索引[0,0,1,1]对应的元素,得到临时数组[0,0,0,0] - 临时数组加上
[1,2,3,4]得到[1,2,3,4] - 把临时数组的结果赋值回原数组的对应索引,此时索引0被连续赋值两次(先1再2),最终保留最后一次的2;索引1同理保留4。
内容的提问来源于stack exchange,提问作者Michael Ma
相关产品推荐
相关产品推荐

