NumPy高级索引更新数组遇问题:+=未按预期累加求解
问题原因与解决方案
问题根源
使用整数数组索引时,foo[bar]返回的是原数组的副本而非视图。执行foo[bar] += barr时,实际流程是:
- 先创建
foo[bar]的副本; - 给副本中的元素加上
barr对应的值; - 将副本的值写回原数组
foo的对应索引位置。
由于bar中索引0出现两次,写回时后一次的赋值(12)会覆盖前一次(3),导致foo[0]最终只保留了最后一次的结果,而非两次累加的总和。
解决方法
方法1:使用np.add.at(推荐)
np.add.at是NumPy专门用于处理重复索引原地累加的函数,能确保每个索引对应的所有值都被正确累加:
import numpy as np foo = np.array([0, 0, 0]) bar = np.array([0, 2, 1, 0]) barr = np.array([3, 6, 9, 12]) np.add.at(foo, bar, barr) print(foo) # 输出: [15 9 6]
方法2:使用np.bincount
通过np.bincount统计每个索引对应的累加权重,直接生成结果数组:
import numpy as np bar = np.array([0, 2, 1, 0]) barr = np.array([3, 6, 9, 12]) foo = np.bincount(bar, weights=barr, minlength=3) print(foo) # 输出: [15 9 6]
内容的提问来源于stack exchange,提问作者willfung
相关产品推荐
相关产品推荐

